diff --git a/.env.example b/.env.example index 35a171ba..265ca2c4 100644 --- a/.env.example +++ b/.env.example @@ -33,6 +33,8 @@ ROUTSTR_SECRET_KEY= # DATABASE_POOL_PRE_PING=false # Warn when a checkout is held this many seconds. # DATABASE_POOL_HOLD_WARN_SECONDS=10 +# SQLite write-lock timeout, in seconds. +# DATABASE_BUSY_TIMEOUT=30 # SQLite serialises writes; increasing its pool can trade pool timeouts for # "database is locked" errors rather than increasing write throughput. diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index dd3ad5a4..b0600ba3 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -11,7 +11,7 @@ jobs: runs-on: ubuntu-latest strategy: matrix: - python-version: ["3.11", "3.12"] + python-version: ["3.11", "3.12", "3.14"] steps: - name: Checkout code @@ -25,22 +25,22 @@ jobs: - name: Install dependencies run: | - uv sync --dev + uv sync --python ${{ matrix.python-version }} --dev - name: Run linting with ruff run: | - uv run ruff check . + uv run --python ${{ matrix.python-version }} ruff check . - name: Run type checking with mypy run: | - uv run mypy . + uv run --python ${{ matrix.python-version }} mypy . - name: Run tests with pytest env: UPSTREAM_BASE_URL: "http://test" UPSTREAM_API_KEY: "test" run: | - uv run pytest --verbose --tb=short + uv run --python ${{ matrix.python-version }} pytest --verbose --tb=short - name: Upload test results if: always() diff --git a/.gitignore b/.gitignore index 903d526e..4b74d8db 100644 --- a/.gitignore +++ b/.gitignore @@ -11,6 +11,9 @@ dist/ *.egg .mypy_cache/** +# MkDocs build output +site/ + # Development .notes .*keys.db diff --git a/.python-version b/.python-version index 2c073331..6324d401 100644 --- a/.python-version +++ b/.python-version @@ -1 +1 @@ -3.11 +3.14 diff --git a/Dockerfile b/Dockerfile index 23140ccf..e4506e40 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,10 +1,12 @@ -FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim +ARG PYTHON_VERSION=3.14 +FROM ghcr.io/astral-sh/uv:python${PYTHON_VERSION}-bookworm-slim RUN apt-get update \ && apt-get install -y --no-install-recommends \ git \ build-essential \ pkg-config \ + libffi-dev \ libsecp256k1-dev \ autoconf \ automake \ diff --git a/Dockerfile.full b/Dockerfile.full index 56406858..a959801f 100644 --- a/Dockerfile.full +++ b/Dockerfile.full @@ -1,4 +1,6 @@ # Multi-stage Dockerfile for Routstr (includes UI build) +ARG PYTHON_VERSION=3.14 + # Stage 1: Build the UI FROM node:23-alpine AS ui-builder WORKDIR /app/ui @@ -16,13 +18,14 @@ ENV NEXT_TELEMETRY_DISABLED=1 RUN pnpm run build # Stage 2: Build the Routstr Node -FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim AS runner +FROM ghcr.io/astral-sh/uv:python${PYTHON_VERSION}-bookworm-slim AS runner RUN apt-get update \ && apt-get install -y --no-install-recommends \ git \ build-essential \ pkg-config \ + libffi-dev \ libsecp256k1-dev \ autoconf \ automake \ diff --git a/docs/api/authentication.md b/docs/api/authentication.md index c7b91b3b..387481cf 100644 --- a/docs/api/authentication.md +++ b/docs/api/authentication.md @@ -307,23 +307,6 @@ ANALYTICS_KEY = os.getenv("ROUTSTR_ANALYTICS_KEY") api_key = PROD_KEY if is_production() else DEV_KEY ``` -### Delegated Authentication - -Create sub-keys with limited permissions: - -```bash -POST /v1/wallet/create/subkey -Authorization: Bearer sk-parent-key -Content-Type: application/json - -{ - "name": "Limited Subkey", - "balance_limit": 1000, - "allowed_models": ["gpt-3.5-turbo"], - "expires_in_hours": 24 -} -``` - ## Rate Limiting Rate limits are applied per API key: @@ -388,22 +371,6 @@ Content-Type: application/json ## Monitoring -### Usage Alerts - -Set up usage notifications: - -```bash -POST /v1/wallet/alerts -Authorization: Bearer sk-... -Content-Type: application/json - -{ - "low_balance_threshold": 1000, - "daily_spend_limit": 5000, - "webhook_url": "https://your-app.com/webhook" -} -``` - ### Audit Logging All API key usage is logged: diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 23b019c7..8e524c72 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -418,7 +418,7 @@ POST /v1/wallet/create ### Get Key Information -Get current balance, consumption data, and child keys for an API key. +Get current balance and consumption data for an API key. ```http GET /v1/balance/info @@ -432,23 +432,9 @@ Authorization: Bearer sk-... "api_key": "sk-abc...", "balance": 8500000, "reserved": 0, - "is_child": false, - "parent_key": null, "total_requests": 42, "total_spent": 1500000, - "balance_limit": null, - "balance_limit_reset": null, - "validity_date": null, - "child_keys": [ - { - "api_key": "sk-child1...", - "total_requests": 10, - "total_spent": 500000, - "balance_limit": 1000000, - "balance_limit_reset": "daily", - "validity_date": 1738000000 - } - ] + "validity_date": null } ``` @@ -500,48 +486,23 @@ Authorization: Bearer sk-... } ``` -### Withdraw Funds +### Refund Balance -Withdraw balance as eCash. +Pay out the current balance. The key remains valid at zero balance and can be topped up again. The payout goes to a Lightning address when one is given (in the request or stored on the key), otherwise a Cashu token is returned. ```http -POST /v1/wallet/withdraw +POST /v1/balance/refund Authorization: Bearer sk-... +Content-Type: application/json ``` -**Request Body:** +`/v1/wallet/refund` is a deprecated alias. + +**Request Body** (optional): ```json { - "amount": 5000, - "mint": "https://mint.example.com" -} -``` - -**Response:** - -```json -{ - "cashu_token": "cashuAeyJ0...", - "amount": 5000, - "mint": "https://mint.example.com" -} -``` - -### Create Child Key - -Creates one or more child API keys that share the parent's balance. Each child key creation costs a fixed amount (configurable). - -```http -POST /v1/balance/child-key -Authorization: Bearer sk-... -``` - -**Request Body:** - -```json -{ - "count": 1 + "lightning_address": "user@getalby.com" } ``` @@ -549,21 +510,67 @@ Authorization: Bearer sk-... | Parameter | Type | Required | Default | Description | |-----------|------|----------|---------|-------------| -| `count` | integer | Yes | - | Number of child keys to create (1-50) | +| `lightning_address` | string | No | Key's stored refund address | Lightning address or LNURL to pay. Overrides the stored address for this request. The effective address (request or stored) is resolved only for a request that can open a new claim, before any balance is debited. | -**Response:** +**Response (Lightning):** ```json { - "api_keys": ["sk-abc...", "sk-def..."], - "count": 2, - "cost_msats": 2000, - "cost_sats": 2, - "parent_balance": 98000, - "parent_balance_sats": 98 + "refund_id": "3f9c1e2d8b7a4c6e9f0a1b2c3d4e5f60", + "status": "paid", + "recipient": "user@getalby.com", + "sats": "4500" } ``` +**Response (Cashu):** + +```json +{ + "refund_id": "3f9c1e2d8b7a4c6e9f0a1b2c3d4e5f60", + "status": "paid", + "token": "cashuAeyJ0...", + "sats": "4500" +} +``` + +The amount field is `sats` or `msats` depending on the key's refund currency. It reports the gross balance debited by the claim. For Lightning refunds, mint and input fees can reduce the amount actually delivered to the recipient. + +**Behaviour:** + +- The balance is debited and a refund claim is recorded before the payout is attempted. A key has at most one open claim at a time. +- If the payout fails cleanly, the claim is closed and the balance is restored. Retry the request. +- Once a melt quote has been recorded or a Cashu token has been issued, the payout may already have happened, so any later failure returns `502` and withholds the balance rather than restoring it. The exception is the mint answering the melt itself with `unpaid`: that is proof nothing was sent, so the balance is restored at once and the request returns `503`. +- If the Lightning payment is dispatched but the mint cannot confirm the outcome, the request returns `502`, the balance stays withheld, and a background reconciler asks the mint until it answers. The balance is restored if the mint reports the payment unpaid. +- An unresolved claim is reported before any replay: a request on a key with an open claim returns `409` with that claim's `refund_id` and `status`. +- Calling again on a zero-balance key with no open claim returns the last paid Lightning refund, or the Cashu token issued by the last paid claim while it remains uncollected. + +**Errors:** + +| Status | Meaning | +|--------|---------| +| `400` | Invalid Lightning destination, no balance, or balance too small for the refund unit | +| `400` | Ongoing requests are still reserving balance on this key | +| `401` | Unknown key | +| `409` | Balance changed concurrently. Retry. | +| `409` | `refund_in_progress`: another refund claim for this key is still open. The body carries its `refund_id` and `status` | +| `409` | `refund_unresolved`: a claim for this key is `stuck` and needs operator reconciliation | +| `410` | Previously issued Cashu refund token has been swept | +| `500` | Payout failed before anything was dispatched. Balance restored. Retry. | +| `502` | Payment dispatched, outcome unconfirmed. Balance withheld pending reconciliation. Do not retry. | +| `503` | Mint unavailable, or the mint reported the Lightning payment unpaid. Balance restored. Retry later. | + +**X-Cashu refunds:** + +Requests paid per-call with an `X-Cashu` header get their change from this endpoint by sending the same header instead of `Authorization`: + +```http +POST /v1/balance/refund +X-Cashu: cashuAeyJ0... +``` + +Returns the change token in the body and in an `X-Cashu` response header. `404` if no matching request exists, `425` while the change is still being minted, `410` if it was swept. + ## Provider Discovery ## Admin Settings diff --git a/docs/api/errors.md b/docs/api/errors.md index e74a46f4..82bb5be1 100644 --- a/docs/api/errors.md +++ b/docs/api/errors.md @@ -153,28 +153,32 @@ granularity) on any of them. |--------|--------|--------|-----------|---------| | `token_already_spent` | 400 | `cashu_token_already_spent` | No | The token was already redeemed. | | `invalid_token` | 400 | `invalid_cashu_token` | No | The token is malformed or cannot be decoded. | -| `mint_error` | 422 | `cashu_token_swap_fees_exceed_amount` | No | Token value is too small to cover the mint's swap/melt fees. | -| `mint_error` | 422 | `cashu_foreign_mint_swap_failed` | No | Swapping the token from a foreign mint to the primary mint failed. | +| `mint_error` | 422 | `cashu_token_swap_fees_exceed_amount` | No | Token value is too small to cover the mint's NUT-02 input fees. | +| `untrusted_mint` | 400 | `cashu_untrusted_source_mint` | No | The token was issued by a mint this node does not accept. Only the node's configured mints (`PRIMARY_MINT_URL` / `CASHU_MINTS`) are redeemable. | | `mint_unreachable` | 503 | `cashu_source_mint_unreachable` | **Yes** | The mint that issued the token could not be reached; it cannot be redeemed at another mint. | | `mint_rate_limited` | 503 | `cashu_mint_rate_limited` | **Yes** | The mint rate-limited the request; retry after the cooldown. | -| `mint_unreachable` | 503 | `cashu_mint_unreachable` | **Yes** | The mint could not be reached (DNS failure, refused/reset connection, timeout). The token is fine — retry once the mint recovers. | +| `mint_timeout` | 503 | `cashu_mint_timeout` | **Yes** | The mint did not respond in time; retry later. | +| `mint_unreachable` | 503 | `cashu_mint_unreachable` | **Yes** | The mint could not be reached (DNS failure, refused/reset connection). The token is fine — retry once the mint recovers. | | `cashu_error` | 400 | `cashu_token_redemption_failed` | No | The token could not be redeemed for another expected reason. | | `cashu_error` | 400 | `cashu_token_zero_value` | No | The token redeemed to zero (empty/dust token, or value fully consumed by fees). | | `token_consumed` | 500 | `cashu_token_consumed` | No | The token was **spent** (melted/redeemed) but crediting it then failed. Do not retry — the token is gone; contact support to reconcile. | | `api_error` | 500 | `internal_error` | Maybe | Unexpected server-side fault during redemption. | !!! important "Retry only transient mint failures" - Only `mint_unreachable` and `mint_rate_limited` (503) are retryable — the - same token may work again later. Everything else is a permanent property of - the token and must not be blindly retried. Use exponential backoff for the + Only `mint_unreachable`, `mint_rate_limited` and `mint_timeout` (503) are + retryable — the same token may work again later. Everything else is a + permanent property of the token and must not be blindly retried. + `untrusted_mint` is permanent: the node will never accept that mint until + an operator adds it to `CASHU_MINTS`. Use exponential backoff for the 503 responses, and honor the mint's cooldown for `mint_rate_limited`. In particular, a `token_consumed` 500 means the mint already spent the token, so a retry would fail as `token_already_spent`. #### Mint failures (retryable) -`mint_unreachable` and `mint_rate_limited` are retryable redemption errors. For -`mint_rate_limited`, honor the mint's cooldown before retrying. +`mint_unreachable`, `mint_rate_limited` and `mint_timeout` are retryable +redemption errors. For `mint_rate_limited`, honor the mint's cooldown before +retrying. ```json { diff --git a/docs/client/payments.md b/docs/client/payments.md index f85c8df2..cd118b13 100644 --- a/docs/client/payments.md +++ b/docs/client/payments.md @@ -40,8 +40,9 @@ Every time you make a request to `/v1/chat/completions` (or others), the cost is `Cost = (Input_Tokens * Price_Input) + (Output_Tokens * Price_Output) + Request_Fee` - Prices are defined per model (see `/v1/models`). -- If you stream the response, the balance is deducted incrementally or finalized at the end of the stream. -- If your balance hits 0 mid-stream, the connection is closed. +- Routstr reserves an authorization ceiling before forwarding, then finalizes the request at measured token cost. +- If a successful upstream omits usage, Routstr estimates input tokens from the provider-bound request and output tokens from the returned body or streamed deltas, then applies normal model pricing. +- A reservation is only a temporary hold. Missing usage or unusable prices must never turn the full reservation into the charge; if no auditable estimate can be priced, the reservation is released without charge. ### Headers diff --git a/docs/contributing/code-structure.md b/docs/contributing/code-structure.md index 40432164..b73e2536 100644 --- a/docs/contributing/code-structure.md +++ b/docs/contributing/code-structure.md @@ -300,7 +300,7 @@ Project metadata and dependencies: name = "routstr" version = "0.2.2" dependencies = [ - "fastapi[standard]>=0.115", + "fastapi[standard-no-fastapi-cloud-cli]>=0.141", "sqlmodel>=0.0.24", "cashu", # ... diff --git a/docs/ehbp-proxy-support.md b/docs/ehbp-proxy-support.md index f9edc6e0..c5bc89d1 100644 --- a/docs/ehbp-proxy-support.md +++ b/docs/ehbp-proxy-support.md @@ -58,10 +58,10 @@ Contains the shared opaque EHBP transport and billing helpers: - `EHBPForwardingTarget` — provider-specific target URL plus extra headers - `forward_ehbp_request()` — forwards the encrypted body, captures Tinfoil usage from a response header or streaming HTTP trailer, and finalizes bearer - billing at actual cost (falling back to max cost when usage is unavailable) + billing at actual cost (releasing the reservation when usage is unavailable) - `forward_ehbp_x_cashu_request()` — redeems the Cashu token, refunds the full token on upstream failure, and refunds the difference between the redeemed - amount and actual cost (or max cost when usage is unavailable) + amount and actual cost (or the full amount when usage is unavailable) ### Provider support @@ -84,7 +84,7 @@ The proxy is a **blind relay** for EHBP requests. It cannot decrypt the body Cost tracking happens at the proxy level. Routstr reserves or redeems up to `max_cost_for_model`, then Tinfoil's out-of-band usage header/trailer allows it to finalize at actual token cost. If trusted usage is missing or invalid, the -proxy safely falls back to max-cost billing. +proxy releases/refunds rather than treating the authorization ceiling as usage. ## End-to-end flow diff --git a/docs/index.md b/docs/index.md index 9a72b9b7..6615eaee 100644 --- a/docs/index.md +++ b/docs/index.md @@ -26,6 +26,20 @@ If you want to run a node, resell API access, or monetize hardware. --- +## 👥 For Teams (Remote Nodes) + +If you want one shared Routstr endpoint for a whole team, with per-member identities and per-member usage tracking. + +- **[Overview](teams/index.md)**: What a remote node is, and when to use one. +- **[Deploy on Cloudron](teams/deploy-cloudron.md)**: The packaged, supported deployment. +- **[Deploy with Docker](teams/deploy-docker.md)**: Run it on any host behind your own TLS. +- **[Team Members](teams/team-members.md)**: Bootstrap the first admin and invite people. +- **[Connecting Clients](teams/clients.md)**: Wire up Claude Code, Pi, OpenCode, and API keys. +- **[Usage and Model Policy](teams/usage-and-policy.md)**: Per-member spend and model allowlists. +- **[Security Model](teams/security.md)**: Auth rules and endpoint scoping. + +--- + ## 🔌 API Reference - **[Overview](api/overview.md)**: Base URL, headers, and standards. diff --git a/docs/provider/discovery.md b/docs/provider/discovery.md index c0d06036..63e6caaf 100644 --- a/docs/provider/discovery.md +++ b/docs/provider/discovery.md @@ -9,6 +9,19 @@ Routstr uses **Nostr** as a decentralized directory for service discovery. Your 1. **Provider Advertisement (Kind 38421)**: Your node periodically publishes an event with its URL, models, and pricing 2. **Client Discovery**: Clients query relays for these events to find suitable providers +### When announcements are published + +Your node publishes an advertisement as soon as it has both a **Nsec** and at least one +reachable endpoint (a public `HTTP_URL`, or an `.onion` address). Saving the Nsec in the +dashboard is enough — the announcement follows within a minute, and **no restart is +required**. After the first publish it re-announces every 24 hours, and immediately +whenever the Nsec, endpoints, mints or relays change. + +A node that has no Nsec yet simply waits, and starts announcing the moment one is +configured. Note that `HTTP_URL` defaults to `http://localhost:8000`, which is not a +reachable endpoint: a node with the default value and no onion address has nothing to +advertise and will not publish until one is set. + --- ## Configuration diff --git a/docs/teams/clients.md b/docs/teams/clients.md new file mode 100644 index 00000000..1ec24a13 --- /dev/null +++ b/docs/teams/clients.md @@ -0,0 +1,153 @@ +# Connecting Clients + +A **client** is one agent or application talking to the node. Each client gets its own ID and its own API key (`sk-...`). Clients are how a team node attributes usage: every request is billed against the client that made it. + +Members create and manage **their own** clients. They cannot see or touch anyone else's. + +--- + +## Two credentials, two jobs + +The node accepts two entirely different kinds of credential, and confusing them is the most common source of `403`s. + +| Credential | Header | Purpose | Can do | +|---|---|---|---| +| **API key** | `Authorization: Bearer sk-...` | Inference | Send chat/completion requests. Cannot touch wallets, clients, or npubs. | +| **NIP-98** | `Authorization: Nostr ` | Management | Manage clients, npubs, wallet, node control. Signed per-request by the member's `nsec`. | + +Your agents use the **API key**. The `routstrd` CLI uses **NIP-98** automatically, which is why it needs your `nsec` in `~/.routstrd/config.json`. + +--- + +## Add a client + +From the member's own machine: + +```bash +# Name it explicitly +routstrd clients add --name "My Laptop" + +# Or use a one-shot integration setup +routstrd clients add --claude-code +routstrd clients add --pi-agent +routstrd clients add --opencode +routstrd clients add --openclaw +routstrd clients add --hermes +``` + +The integration flags configure the agent's own config file as well as registering the client, so you do not have to hand-edit anything. Several can be combined in one call. + +On success the CLI prints the credentials and the endpoint to point at: + +```text +Client created. + + ID: my-laptop + Name: My Laptop + API Key: sk-9f2a... + + Access Routstr at: https://team.example.com/v1 +``` + +!!! warning "The API key is a secret" + Treat `sk-...` like a password. It bills inference to the team wallet. Do not commit it, and do not paste it into a chat — unlike an npub, it is not safe to share. + +### Adding is idempotent + +Running `clients add` with a name that already exists does not create a duplicate. It looks the client up and prints the existing record — including its API key — so re-running is a safe way to recover a key you lost: + +```text +Client 'my-laptop' already exists. + + ID: my-laptop + Name: My Laptop + API Key: sk-9f2a... +``` + +### List and delete + +```bash +routstrd clients list +routstrd clients delete my-laptop +``` + +`clients list` shows **only your own** clients. + +--- + +## Going beyond the CLI + +The agent integrations cover the common tools, but any OpenAI-compatible client works — point it at the node and use the API key: + +```bash +curl https://team.example.com/v1/chat/completions \ + -H "Authorization: Bearer sk-9f2a..." \ + -H "Content-Type: application/json" \ + -d '{ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}] + }' +``` + +The base URL is always the node host plus `/v1`. Discover available models without any credential at all: + +```bash +curl https://team.example.com/v1/models +``` + +--- + +## How client ownership works + +This is the part that explains the odd-looking IDs on the node. + +### IDs are derived from the name + +A client's ID is the name lowercased with internal whitespace collapsed to hyphens, then stripped of anything that is not alphanumeric or a hyphen. `"My Laptop!"` becomes `my-laptop`. + +### The node appends an owner suffix + +So that two members can both have a client called `my-laptop` without colliding, the auth proxy appends the **last 7 characters of the owner's npub** to the ID before it reaches the daemon: + +| Where you look | Client ID | +|---|---| +| Member's `routstrd clients list` | `my-laptop` | +| On the node (`cloudron exec`, then `routstrd clients list`) | `my-laptop-4f2x9k7` | + +The suffix is stripped again on the way back, so members always see the clean ID. On the node you deliberately see the suffixed form — **the trailing characters are what tell you which member owns a client.** + +### Ownership is recorded explicitly + +Newly created clients store the owner's npub in an `ownerNpub` field, and the proxy authorises against that field. Clients created before that field existed fall back to matching the ID suffix, so older installs keep working until those clients are recreated. + +### Consequence: admins are not automatically superusers here + +`/clients`, `/clients/add`, and `/clients/delete` are owner-scoped by the calling npub. Even an `admin` cannot list or delete a colleague's clients through these endpoints. Cross-member visibility comes from running the CLI **on the node itself**, where the daemon is unauthenticated on loopback: + +```bash +cloudron exec --app routstr.example.com +routstrd clients list # all clients, all owners, suffixed IDs +``` + +That is also the only practical way to clean up a departing member's keys — see [Team Members](team-members.md#what-revocation-does-and-does-not-do). + +--- + +## Refreshing models and integrations + +The `clients` command carries options for the daemon's scheduled refresh job, which updates the Routstr 21 model list and re-syncs client integrations: + +```bash +routstrd clients --manual-refresh # refresh now, once +routstrd clients --disable-automatic-refresh # stop the scheduled job +routstrd clients --enable-automatic-refresh # start it again +``` + +The model list matters because it is also what the [model allowlist](usage-and-policy.md#model-allowlist) is enforced against. + +--- + +## Next steps + +- [Usage and Model Policy](usage-and-policy.md) — watch what those clients are spending. +- [Security Model](security.md) — the exact rules applied to each credential. diff --git a/docs/teams/deploy-cloudron.md b/docs/teams/deploy-cloudron.md new file mode 100644 index 00000000..9f1a7805 --- /dev/null +++ b/docs/teams/deploy-cloudron.md @@ -0,0 +1,169 @@ +# Deploy on Cloudron + +[Cloudron](https://www.cloudron.io/) is the supported deployment target for a team node. The packaged image already contains **both** processes — the `routstrd` daemon and the `routstrd-auth` proxy — supervised inside a single container, with `/app/data` handled as persistent storage and TLS terminated by the platform. + +| | | +|---|---| +| **App ID** | `io.routstr.routstrd-auth` | +| **Public port** | `8008` (Cloudron proxies it over HTTPS on 443) | +| **Health check** | `GET /health` | +| **Memory limit** | 512 MB | +| **Minimum box version** | Cloudron 9.1.0 | + +--- + +## Prerequisites + +- A running Cloudron box with a domain that can get a certificate. +- The [`cloudron` CLI](https://docs.cloudron.io/cli/) installed and logged in, **only if** you are building the image yourself: + +```bash +npm install -g cloudron +cloudron login my.example.com +``` + +## Install + +### Option A — from the published version list + +The app is published as a custom Cloudron app with a version list (`CloudronVersions.json`), currently at `0.1.26`. Once that app store entry is registered on your Cloudron instance, install it from the dashboard, or: + +```bash +cloudron install --appstore-id io.routstr.routstrd-auth --location routstr.example.com +``` + +### Option B — build the image yourself + +Use this when you want to run a local modification: + +```bash +git clone https://github.com/routstr/routstrd-remote +cd routstrd-remote + +cloudron build # builds the Dockerfile and pushes it to your registry +cloudron install --image /routstrd-remote: --location routstr.example.com +``` + +!!! note "The Dockerfile is the Cloudron image" + The repository's `Dockerfile` is built `FROM cloudron/base:5.0.0` and its `CMD` is `cloudron/start.sh`, which prepares `/app/data` and starts `supervisord`. It expects Cloudron's filesystem conventions and should not be confused with a generic Docker image. See [Deploy with Docker](deploy-docker.md) for what that means in practice. + +--- + +## Bootstrap the first admin + +The moment the app is healthy, the npub table is **empty**, and nothing except the public endpoints can be reached. Claim it before anything else — while the table is empty, `POST /npubs` is accepted without authentication, so this is the only window in which an unauthenticated registration succeeds. + +On the machine of whoever will be the first admin: + +```bash +bun i -g routstrd +routstrd remote https://routstr.example.com +routstrd npubs register --name "Alice" +``` + +`routstrd remote` generates a fresh Nostr identity if you do not have one, stores it in `~/.routstrd/config.json`, and prints your npub. `routstrd npubs register` then posts that npub and, because no npubs exist yet, receives `admin`. + +!!! warning "Register immediately after install" + Until the first admin registers, anyone who knows the URL can claim the node. Do this as part of the install, not later. + +Verify: + +```bash +routstrd npubs list +``` + +--- + +## Configuration + +Cloudron defaults are set by `cloudron/start.sh` and the two supervisor programs. Everything below can be overridden through the Cloudron **Environment Variables** tab. + +| Variable | Default | Purpose | +|---|---|---| +| `ROUTSTRD_AUTH_PORT` | `8008` | Public port served by the auth proxy. Must match the manifest's `httpPort`. | +| `ROUTSTRD_AUTH_HOST` | `0.0.0.0` | Bind address of the auth proxy. | +| `ROUTSTRD_UPSTREAM` | `http://localhost:8009` | Where the daemon listens. Keep this on loopback. | +| `ROUTSTRD_PORT` | `8009` | Port the daemon binds. | +| `ROUTSTRD_DIR` | `/app/data/routstrd` | Config directory shared by the daemon and the proxy. | +| `ROUTSTRD_DB_PATH` | `/app/data/routstrd/routstr.db` | Shared SQLite database. | +| `ROUTSTRD_CONFIG_FILE` | `$ROUTSTRD_DIR/config.json` | Daemon config file. | +| `ROUTSTRD_AUTH_MODEL_ALLOWLIST` | `false` | Set to `true` to restrict the team to the Routstr 21 model list. See [Usage and Model Policy](usage-and-policy.md). | +| `ROUTSTRD_AUTH_ADMIN_NPUBS` | *(unset)* | Optional bootstrap admins. See below. | + +### Bootstrapping admins from the environment + +Instead of the interactive `npubs register` step you can seed admins declaratively. Three variables are accepted and merged: `ROUTSTRD_AUTH_ADMIN_NPUBS`, `ROUTSTRD_AUTH_ADMIN_PUBKEYS`, and `ROUTSTRD_AUTH_BOOTSTRAP_NPUB`. Values are comma- or whitespace-separated and may be either `npub1...` or 64-character hex. + +Rows created this way are tagged `source = 'env'`. At every startup the proxy **reconciles** them: an env-sourced row whose pubkey is no longer present in the environment is **deleted**. This means the environment variables are the source of truth for those rows — removing someone from the variable revokes their access on the next restart. + +!!! tip "Prefer `npubs register` for the first admin" + There is deliberately **no** hardcoded default admin npub in the image. An image with a baked-in admin pubkey would hand control of every deployment to the same key. + +### Filesystem layout + +| Path | Lifetime | Contents | +|---|---|---| +| `/app/code` | replaced on update | Auth proxy source and the `start.sh` / `run-auth.sh` scripts. | +| `/app/data` | persistent, backed up | `routstrd/config.json`, `routstrd/routstr.db`, `logs/`, and a `.initialized` marker. | +| `/run` | ephemeral | `supervisord` socket and pid file. | + +The startup script writes `authUrl` into `routstrd/config.json` pointing at the local proxy, and generates the container's Nostr identity (`nsec`) on first boot if one is missing. That generated identity is what authorises the daemon's own calls back through the proxy. + +--- + +## Backups + +Cloudron's `localstorage` addon makes `/app/data` persistent and includes it in regular backups. The entire data directory is covered, which means the SQLite database at `/app/data/routstrd/routstr.db` travels with it. + +To restore, restore the app from a Cloudron backup — the wallet configuration, npub table, and client records all come back together. Do not hand-copy files between hosts: the daemon's `config.json` holds the container's `nsec`, and losing it breaks the node's ability to authenticate to its own proxy. + +!!! warning "Use `cloudron exec` carefully" + The database uses SQLite WAL mode. If you need a manual snapshot, stop the app first (`cloudron stop`) or use `sqlite3 .backup` inside the container — copying the `.db` file while the daemon is writing can produce an inconsistent file. + +--- + +## Updates + +```bash +cloudron update --app routstr.example.com +``` + +Updates are the primary lifecycle event on Cloudron and are designed to preserve `/app/data`. After the app comes back, confirm health and that your admin npub is still recognised: + +```bash +curl https://routstr.example.com/health +routstrd npubs list +``` + +## Day-to-day operations + +```bash +cloudron logs -f --app routstr.example.com # follow both processes' output +cloudron exec --app routstr.example.com # shell into the container +cloudron stop --app routstr.example.com +cloudron start --app routstr.example.com +cloudron debug --app routstr.example.com # read-write filesystem, app paused +cloudron debug --disable --app routstr.example.com +``` + +Inside the container you are on the machine that holds the wallet and the shared database, so `routstrd` commands there operate on the node itself rather than as a remote member: + +```bash +routstrd npubs list # everyone with access, and their roles +routstrd clients list # every client on the node, not just your own +routstrd top # interactive usage TUI across all members +``` + +Both processes log to stdout/stderr and are collected by Cloudron; there are no log files to rotate inside `/app/data`. + +### Failure handling + +`supervisord` runs both programs with `autorestart=true` and a start priority that brings the **daemon up first** (priority 10) and the **proxy second** (priority 20). The proxy's launcher additionally waits — up to 120 seconds — for the database file to exist *and* for the daemon's `/health` to answer before it starts serving. A crash-looping proxy therefore usually means the daemon never became healthy. + +--- + +## Next steps + +- [Team Members](team-members.md) — invite the rest of your team. +- [Security Model](security.md) — what is exposed and what is not. +- [Troubleshooting](troubleshooting.md) — when the proxy will not start. diff --git a/docs/teams/deploy-docker.md b/docs/teams/deploy-docker.md new file mode 100644 index 00000000..93f3fea2 --- /dev/null +++ b/docs/teams/deploy-docker.md @@ -0,0 +1,170 @@ +# Deploy with Docker + +The team node is a single container running two supervised processes. You can run that same image on any Docker host and terminate TLS with your own reverse proxy. + +!!! note "Read this first" + The image is built `FROM cloudron/base:5.0.0` and its entrypoint is the Cloudron startup script. It works outside Cloudron — it only requires a writable `/app/data` — but its filesystem conventions are Cloudron's, and the repository's `docker-compose.yml` predates this image. The commands below are the ones that match the current image. + +--- + +## Build + +```bash +git clone https://github.com/routstr/routstrd-remote +cd routstrd-remote +docker build -t routstr-remote:0.1.26 . +``` + +The Dockerfile installs Bun (the x64 **baseline** build, chosen because some hosts do not expose AVX/AVX2 and the default binary crashes with `SIGILL`), installs the `routstrd` daemon globally, installs the proxy's dependencies, and copies in the supervisor configuration for both processes. + +## Run + +```bash +mkdir -p "$HOME/routstr-remote-data" + +docker run -d \ + --name routstr-remote \ + --restart unless-stopped \ + -p 127.0.0.1:8008:8008 \ + -v "$HOME/routstr-remote-data:/app/data" \ + --memory 1g \ + routstr-remote:0.1.26 +``` + +**Why each flag matters:** + +| Flag | Reason | +|---|---| +| `-v ...:/app/data` | The only persistent path. It holds `routstrd/config.json` (including the container's `nsec`), `routstrd/routstr.db`, and `logs/`. **Without it the node loses its identity and every npub on restart.** | +| `-p 127.0.0.1:8008:8008` | Publish on loopback only and let a TLS reverse proxy in front of it expose the service. Binding `0.0.0.0:8008` on a public host sends API keys over plaintext HTTP. | +| `--memory 1g` | Two Bun processes plus the daemon's model and usage state. The Cloudron manifest requests 512 MB; give a bare Docker host at least that, and prefer more. | + +The container listens on exactly two ports: `8008` (public auth proxy) and `8009` (daemon, bound to loopback **inside** the container and deliberately not published). + +Check it came up: + +```bash +curl http://127.0.0.1:8008/health +docker logs -f routstr-remote +``` + +### Startup ordering + +The entrypoint starts `supervisord`, which brings up the daemon first and the proxy second. The proxy's launcher then waits for the database file to appear and for `http://localhost:8009/health` to answer, retrying up to 120 times at one-second intervals. If that window expires, the proxy exits with `Timed out waiting for routstrd to become ready.` and restarts. + +## Put TLS in front + +Anything that terminates TLS and forwards to `127.0.0.1:8008` works. Two properties are worth configuring explicitly: + +- **Disable response buffering.** The proxy already sends `X-Accel-Buffering: no` upstream, but your own proxy should also be configured not to buffer, otherwise streamed LLM responses appear to truncate. +- **Raise idle timeouts.** Model responses can be silent for a long time while reasoning or waiting on tools. The proxy disables Bun's per-request idle timeout for exactly this reason, but nginx's default 60-second `proxy_read_timeout` will still cut streams. Set it to something generous. + +```nginx +location / { + proxy_pass http://127.0.0.1:8008; + proxy_http_version 1.1; + proxy_set_header Host $host; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; + + proxy_buffering off; + proxy_read_timeout 3600s; + proxy_send_timeout 3600s; +} +``` + +## Bootstrap the first admin + +Identical to Cloudron — run this from the admin's own machine: + +```bash +bun i -g routstrd +routstrd remote https://routstr.example.com +routstrd npubs register --name "Alice" +``` + +While the npub table is empty, `POST /npubs` needs no authentication, so the first registration wins. See [Team Members](team-members.md). + +## Configuration + +Override any of the variables from [Deploy on Cloudron](deploy-cloudron.md#configuration) with `-e`, for example: + +```bash +docker run -d \ + --name routstr-remote \ + --restart unless-stopped \ + -p 127.0.0.1:8008:8008 \ + -v "$HOME/routstr-remote-data:/app/data" \ + -e ROUTSTRD_AUTH_MODEL_ALLOWLIST=true \ + routstr-remote:0.1.26 +``` + +Inside the container the defaults are `ROUTSTRD_DIR=/app/data/routstrd`, `ROUTSTRD_DB_PATH=/app/data/routstrd/routstr.db`, `ROUTSTRD_UPSTREAM=http://localhost:8009`, `ROUTSTRD_AUTH_HOST=0.0.0.0`, `ROUTSTRD_AUTH_PORT=8008`. + +!!! warning "Use a bind mount, not an anonymous volume" + If you recreate the container (`docker rm` then `docker run`), an unnamed volume is orphaned and the node comes back with a **new** Nostr identity and an **empty** npub table — meaning the next person to hit `POST /npubs` becomes admin. Always mount a known host directory. + +## Backups + +The database is SQLite in WAL mode, so stop the container before copying the data directory, or take a proper snapshot from inside: + +```bash +docker stop routstr-remote +tar czf routstr-remote-$(date +%F).tar.gz -C "$HOME" routstr-remote-data +docker start routstr-remote +``` + +Or, without downtime: + +```bash +docker exec routstr-remote sqlite3 /app/data/routstrd/routstr.db ".backup /app/data/routstrd/backup.db" +docker cp routstr-remote:/app/data/routstrd/backup.db . +``` + +## Updates + +```bash +docker stop routstr-remote && docker rm routstr-remote +git pull +docker build -t routstr-remote:new . +docker run -d ... routstr-remote:new # same -v and -p flags as before +``` + +Because `/app/data` is a bind mount, the node's identity, npub table, and client records survive the swap. That is the entire reason the mount is not optional. + +--- + +## Running from source (evaluation and development) + +You do not need Docker to try the proxy. Point it at any routstrd daemon's database: + +```bash +bun install +bun run src/index.ts validate # checks config and opens the DB +bun run src/index.ts start # binds 0.0.0.0:8008 +``` + +Useful flags: `--port`, `--host`, `--upstream`, `--db-path`. The `validate` subcommand prints the effective configuration and reports how many npubs are registered, split by role — it is the fastest way to confirm the proxy can see the right database: + +```text +Configuration: + Port: 8008 + Host: 0.0.0.0 + Upstream: http://localhost:8009 + DB path: /app/data/routstrd/routstr.db + Bootstrap admin npubs/pubkeys from env: 0 + Model allowlist: disabled + +✅ DB accessible. 3 npub(s) registered (1 admin, 2 user). +``` + +!!! warning "`validate` fails before the daemon has run once" + If the database does not exist, validation stops with `Database not found at ... Make sure routstrd has been initialized`. The proxy shares the daemon's database; it never creates the schema itself. + +--- + +## Next steps + +- [Team Members](team-members.md) — invite the rest of your team. +- [Security Model](security.md) — which endpoints require what. +- [Troubleshooting](troubleshooting.md) — startup and streaming failures. diff --git a/docs/teams/index.md b/docs/teams/index.md new file mode 100644 index 00000000..5d9da8ee --- /dev/null +++ b/docs/teams/index.md @@ -0,0 +1,112 @@ +# Teams and Remote Nodes + +A normal `routstrd` install is a **single-user daemon** running on your own machine. A **remote node** turns that same daemon into a **shared instance for a team**: one server runs the daemon behind an authentication proxy, and every team member gets their own Nostr identity, their own API keys, and their own usage accounting. + +This is the product sometimes called **Routstrd Remote** — the repository is [`Routstr/routstrd-remote`](https://github.com/routstr/routstrd-remote) and the running service is the package `routstrd-auth`. + +--- + +## Who this is for + +- **Small teams and orgs** that want one funded Routstr endpoint instead of one daemon per laptop. +- **Anyone who wants per-person spend visibility** without building a billing system. +- **Self-hosters** who already run [Cloudron](https://www.cloudron.io/) and want a one-click install. + +## What it gives you + +| Capability | Detail | +|---|---| +| **One endpoint, many people** | Members point their coding agents at a single HTTPS URL. No per-machine setup beyond one CLI command. | +| **Per-member identity** | Each person has their own Nostr keypair (npub). Access is granted and revoked by adding or deleting that npub. | +| **Per-member attribution** | Every client registration gets a unique ID, and usage is reported per client, so you can see who is spending what. | +| **Two levels of privilege** | `admin` (manage people, move funds, control the node) and `user` (run inference, manage only their own clients). | +| **Scoped API keys** | Agent API keys can buy inference but cannot touch the wallet or other members' clients. | +| **Model policy** | Optional allowlist restricts the team to approved models. | +| **One wallet to fund** | The team tops up a single node wallet rather than N personal wallets. | + +## What it is not + +- It is **not** multi-tenant SaaS. Everyone shares the node's wallet and upstream provider set. +- It is **not** a billing or chargeback system. Usage is *attributed* per client; invoicing your teammates is up to you. +- It does **not** give each member a separate balance. See [Usage and Model Policy](usage-and-policy.md) for exactly what is tracked. + +--- + +## Architecture + +Two processes run together inside one container. Only one of them is reachable from outside. + +```mermaid +flowchart TD + CLI["routstrd CLI
on a member laptop"] + Agent["Coding agents
Claude Code, Pi, OpenCode"] + App["App holding an sk- API key"] + + TLS["Reverse proxy
TLS termination on 443"] + Proxy["routstrd-auth
0.0.0.0:8008 public"] + Daemon["routstrd daemon
localhost:8009 no auth"] + DB[("routstr.db
shared SQLite")] + Providers["Upstream model providers"] + + CLI -->|https| TLS + Agent -->|https| TLS + App -->|https| TLS + TLS -->|http| Proxy + Proxy -->|forward| Daemon + Proxy -->|npubs and clients| DB + Daemon -->|usage and models| DB + Daemon -->|inference| Providers +``` + +**The security property that matters:** the daemon runs with **no authentication at all** because it is bound to `localhost` and never published. The auth proxy is the only public surface. If you expose port `8009`, you have removed the entire security model. + +### Components + +| Component | Role | Bind | +|---|---|---| +| `routstrd-auth` | Public auth proxy. Validates credentials, enforces roles and model policy, forwards to the daemon. | `0.0.0.0:8008` | +| `routstrd` | The inference daemon. Owns the wallet, providers, clients, and usage records. | `localhost:8009` | +| `routstr.db` | Shared SQLite database. Holds `routstr_auth_npubs`, `clients`, usage rows, and `sdk_storage` (including the Routstr 21 model list). | on disk | +| Reverse proxy | TLS termination and the public hostname. On Cloudron this is managed by the platform. | `443` | + +--- + +## Roles + +Registration lives in a single table, `routstr_auth_npubs`, where each row has a `role` of `admin` or `user`. + +| Capability | `admin` | `user` | +|---|---|---| +| Run inference with own API keys | yes | yes | +| Create and delete **own** clients | yes | yes | +| Read **own** usage | yes | yes | +| List all registered npubs | yes | yes | +| Add / update / delete npubs | yes | no | +| Send funds from the node wallet | yes | no | +| Node control (providers, refunds, stop) | yes | yes | +| Read wallet balance / status | yes | yes | + +`user` is the default role when an admin adds someone. Promote with `routstrd npubs update --role admin`. + +--- + +## Bootstrap order + +A fresh node has an empty npub table. That produces exactly one unaudited window, and only one: + +1. **The first person** runs `routstrd npubs register` against the new node. Because the table is empty, `POST /npubs` is accepted **without authentication**, and the caller becomes `admin`. +2. From that moment on, **every** npub operation requires NIP-98 auth from an existing admin. A second unauthenticated registration is refused with `409` / "already configured". + +Nobody else can self-register. Team members must be added by an admin. See [Team Members](team-members.md). + +--- + +## Next steps + +- **[Deploy on Cloudron](deploy-cloudron.md)** — the supported, packaged deployment with TLS and backups handled for you. +- **[Deploy with Docker](deploy-docker.md)** — run the same image anywhere, behind your own reverse proxy. +- **[Team Members](team-members.md)** — bootstrap the first admin and invite people. +- **[Connecting Clients](clients.md)** — wire up Claude Code, Pi, OpenCode, and raw API keys. +- **[Usage and Model Policy](usage-and-policy.md)** — per-member spend tracking and the model allowlist. +- **[Security Model](security.md)** — the exact auth rules, public paths, and restricted endpoints. +- **[Troubleshooting](troubleshooting.md)** — diagnosing the failures people actually hit. diff --git a/docs/teams/security.md b/docs/teams/security.md new file mode 100644 index 00000000..5865e817 --- /dev/null +++ b/docs/teams/security.md @@ -0,0 +1,172 @@ +# Security Model + +The whole design rests on one property: **the daemon has no authentication because it is never reachable.** `routstrd` binds loopback-only on port `8009`; the auth proxy on `8008` is the single public surface, and it is the only component that makes authorisation decisions. + +!!! danger "Never publish port 8009" + The daemon is unauthenticated by design. Exposing it — or port-forwarding it for debugging, or forgetting to restrict a Docker port mapping to loopback — removes the entire security model at once. Anyone who reaches it can read the wallet, list every member's clients, and spend the team's funds. + +--- + +## Request decision flow + +The proxy is **default-deny**. A request is only forwarded if some rule explicitly allows it. + +```mermaid +flowchart TD + A["request arrives on 8008"] --> B{"management path?
npubs, clients, usage"} + B -->|yes| C["own handler
NIP-98 required"] + B -->|no| D{"GET or HEAD
on a public path?"} + D -->|yes| E["forward, no auth"] + D -->|no| F{"Authorization header?"} + F -->|missing| G["401"] + F -->|"Bearer sk-..."| H{"key found in clients?"} + H -->|no| I["401"] + H -->|yes| J{"restricted path?
wallet, node control"} + J -->|yes| K["403"] + J -->|no| L["forward with header intact"] + F -->|"Nostr event"| M{"valid NIP-98?
url, method, body hash, sig"} + M -->|no| N["401"] + M -->|yes| O{"pubkey registered?
and role sufficient?"} + O -->|no| P["403"] + O -->|yes| Q["forward, header stripped"] +``` + +Two details are easy to miss: + +- **Public means `GET`/`HEAD` only.** The proxy applies the public-path rule only to read methods, because the daemon routes a `POST` to those same paths as a *paid* request. `GET /v1/models` needs no credential; `POST /v1/models` does. +- **Management paths are matched before anything else.** `/npubs`, `/clients`, `/clients/add`, `/clients/delete`, `/usage`, and `/usage/summary` are handled by the proxy's own handlers and never forwarded to the daemon wholesale. + +--- + +## Public paths + +Reachable with no credential at all (`GET`/`HEAD` only): + +| Path | Purpose | +|---|---| +| `/health` | Liveness. Used by Cloudron's health check and by the proxy's own upstream probe. | +| `/ping` | Lightweight reachability. | +| `/models` | Model directory. | +| `/v1/models` | OpenAI-compatible model list. | +| `/models/*`, `/v1/models/*` | Prefixes covering per-model detail paths. | + +This is intentional — an agent needs to discover models before it has a key, and provider discovery is public information in Routstr. + +--- + +## Credential one: API keys (`Bearer sk-...`) + +An API key is looked up in the client records. If no client carries it, the request is rejected with `401 Invalid API key.` + +Keys are **deliberately narrow**. A valid key is refused with `403` on every restricted path: + +| Restricted endpoint | Why | +|---|---| +| `/wallet/status`, `/wallet/unlock`, `/wallet/balance` | An inference key must not read wallet state. | +| `/wallet/receive/cashu`, `/wallet/receive/bolt11` | No minting funds with a key. | +| `/wallet/send/cashu`, `/wallet/send/bolt11` | Admin-only in any case. | +| `/wallet/mints`, `/wallet/mints/info` | No mint inspection. | +| `/stop`, `/refund`, `/refund/xcashu` | No node control. | +| `/providers`, `/providers/enable`, `/providers/disable` | `?refresh=true` rewrites the stored provider list. | +| `/nwc/*` | No payment-channel access. | +| `/npubs`, `/clients/add`, `/clients/delete`, `/usage` | No management surface at all. | + +The rule of thumb: **an API key buys inference, and nothing else.** + +Key handling on the way through: the `Authorization` header is **preserved** so the daemon can validate the key itself, and its own accounting stays authoritative. + +--- + +## Credential two: NIP-98 (`Nostr `) + +Management and wallet operations require a signed [NIP-98](https://github.com/nostr-protocol/nips/blob/master/98.md) event. This is not a bearer token — it is a signature over the specific request, so it cannot be replayed against a different endpoint. + +The proxy enforces, in order: + +| Check | Rule | +|---|---| +| Event kind | must be `27235` | +| Timestamp | within **±60 seconds** of now | +| `u` tag | must equal the **absolute request URL**, including scheme and host | +| `method` tag | must match the HTTP method (case-insensitive) | +| `payload` tag | **required when the body is non-empty**; must equal the SHA-256 hex digest of the raw body, compared in constant time | +| Signature | verified with `verifyEvent` | + +The proxy then looks the pubkey up in `routstr_auth_npubs`: + +- **Not registered** → `403`. The error message is context-aware: on a node with no npubs at all it tells you to run `routstrd npubs register`; otherwise it says registered auth is required. +- **Registered but role insufficient** → `403 Admin access required.` +- **Registered and sufficient** → forwarded, with the `Authorization` header **stripped** so it does not reach the daemon or the upstream provider. + +!!! warning "Behind a reverse proxy, forwarded headers are not optional" + The `u` tag is checked against the **public** URL the client signed. The proxy reconstructs that URL from `X-Forwarded-Proto` and `X-Forwarded-Host` (falling back to `Host`). If your reverse proxy does not set them, the comparison fails and every NIP-98 request is rejected with `NIP-98 URL tag does not match this request.` See the nginx snippet in [Deploy with Docker](deploy-docker.md#put-tls-in-front). + +!!! note "The ±60 second window means clock skew matters" + A client whose clock is more than a minute off will produce events that are rejected as `outside the allowed window`. If one machine alone fails to authenticate, check its clock before suspecting the node. + +--- + +## Role requirements by endpoint + +| Endpoint group | Required | +|---|---| +| `/wallet/send/cashu`, `/wallet/send/bolt11` | `admin` | +| `/wallet/status`, `/wallet/unlock`, `/wallet/balance`, `/wallet/receive/*`, `/wallet/mints*`, `/stop`, `/refund*`, `/providers*`, `/nwc/*` | any registered npub (`admin` or `user`) | +| `/clients`, `/clients/add`, `/clients/delete` | any registered npub, **scoped to own clients** | +| `/usage`, `/usage/summary` | any registered npub, **scoped to own usage** | +| `/npubs` read | any registered npub | +| `/npubs` create / update / delete | `admin` | +| Everything else | valid API key **or** registered npub | + +--- + +## Bootstrap window + +While `routstr_auth_npubs` is empty, `POST /npubs` is accepted **without authentication**. This exists solely so a fresh node can be claimed, and it closes permanently after the first registration. + +The practical implication: a node that is deployed and healthy but has not had its first admin register is **unclaimed**. Treat deployment and bootstrap as one operation. + +There is deliberately no hardcoded default admin in the image. The absence of one means an image cannot be shipped with a known admin key — but it also means a half-finished deployment is claimable by whoever finds it first. + +As an alternative to the interactive step, admins can be seeded with `ROUTSTRD_AUTH_ADMIN_NPUBS`, `ROUTSTRD_AUTH_ADMIN_PUBKEYS`, or `ROUTSTRD_AUTH_BOOTSTRAP_NPUB`. Rows created this way are tagged `source = 'env'` and **reconciled at every startup** — remove the value from the environment and the row is deleted, which is a clean way to authorise a node declaratively. + +--- + +## CORS + +The proxy answers with: + +```text +Access-Control-Allow-Origin: * +Access-Control-Allow-Methods: GET, POST, PATCH, DELETE, OPTIONS +Access-Control-Allow-Headers: Authorization, Content-Type, X-Cashu, X-Routstr-Model +Access-Control-Expose-Headers: X-Cashu, X-Routstr-Request-Id, X-Routstr-Cost-Msats, + X-Routstr-Cost-Usd, X-Routstr-Input-Cost-Msats, + X-Routstr-Output-Cost-Msats +``` + +A wildcard origin is safe **here specifically** because the app uses no cookies and no sessions. There is no ambient browser identity for a cross-origin page to borrow — a malicious page cannot make an authenticated request on a visitor's behalf, because every non-public request still needs its own API key or signature. + +!!! warning "If you ever add cookie or session auth, revisit this" + The wildcard is only correct while authentication is entirely credential-based. Adding session cookies would turn this into a real vulnerability. + +--- + +## Hardening checklist + +- [ ] **Daemon is loopback-only.** Verify `8009` is not published (`docker port routstr-remote`, or check the Cloudron app's port config). +- [ ] **TLS everywhere.** No member or agent should ever send an `sk-...` key over plaintext HTTP. +- [ ] **First admin registered** immediately after install. +- [ ] **Reverse proxy sets `X-Forwarded-Proto` / `X-Forwarded-Host`**, or NIP-98 fails. +- [ ] **Streaming hangs are fixed with timeouts, not by buffering.** Disable `proxy_buffering` and raise `proxy_read_timeout`; do not "fix" a truncated stream by publishing the daemon directly. +- [ ] **`ROUTSTRD_AUTH_ADMIN_NPUBS` reflects reality** if you use env bootstrapping — those rows are deleted on restart when the variable changes. +- [ ] **Departed members have their clients deleted**, not just their npub. Deleting an npub does **not** revoke existing API keys. +- [ ] **Backups cover `/app/data`** and are tested, including the `nsec` in `routstrd/config.json`. +- [ ] **Model allowlist verified** if you rely on it — it fails open when the model list has not been populated. + +--- + +## Next steps + +- [Troubleshooting](troubleshooting.md) — diagnosing `401` and `403` responses. +- [Team Members](team-members.md#what-revocation-does-and-does-not-do) — the two-step offboarding that revocation alone does not cover. diff --git a/docs/teams/team-members.md b/docs/teams/team-members.md new file mode 100644 index 00000000..cb3fd8ba --- /dev/null +++ b/docs/teams/team-members.md @@ -0,0 +1,177 @@ +# Team Members + +Access to a team node is an entry in one table: `routstr_auth_npubs`. Each row holds a Nostr pubkey, an optional display name, and a `role` of `admin` or `user`. Adding someone grants access; deleting their row revokes it. + +There are no passwords, no invite links, and no email addresses. Identity is a Nostr keypair that each person generates on their own machine. + +--- + +## The invite loop + +A new member does the first two steps themselves; an existing admin does the third. + +```mermaid +sequenceDiagram + participant M as New member + participant A as Existing admin + participant N as Team node + + M->>M: install the routstrd CLI + M->>N: set the remote URL + N-->>M: generates keypair and prints npub + M->>A: send npub out of band + A->>N: add npub with role user + N-->>A: access confirmed + M->>N: add a client integration + N-->>M: API key issued +``` + +### 1. The member installs the CLI and connects + +```bash +bun i -g routstrd +routstrd remote https://team.example.com +``` + +`routstrd remote` writes the node URL into `~/.routstrd/config.json` and, **only if no Nostr identity exists yet**, generates one and prints the npub: + +```text +Remote daemon URL set to: https://team.example.com + +A new Nostr identity has been generated for remote authentication. +Your npub: npub1abc...xyz +You can view it in the config file at: /home/bob/.routstrd/config.json +``` + +If you already had an identity, it is reused and no npub is printed — run `routstrd remote` with no arguments to display the current node and identity. + +!!! tip "Your npub is not a secret" + It is a public key and safe to paste into a team chat. The corresponding `nsec` lives in `~/.routstrd/config.json` and **is** a secret: it signs every management request. Never share it, and treat any host that has it as holding that member's credentials. + +### 2. The member sends their npub to an admin + +Out of band — chat, ticket, whatever. There is no self-service join. + +### 3. An admin adds them + +```bash +routstrd npubs add npub1abc...xyz --name "Bob" +``` + +The role defaults to `user`. The new member can now run inference and manage their own clients. To make them an admin, pass `--role admin` (or promote later). + +--- + +## Bootstrap: the first admin + +A brand-new node has an empty table, which is a special case: `POST /npubs` is accepted **without authentication** so that someone can claim the node. + +```bash +routstrd remote https://team.example.com +routstrd npubs register --name "Alice" +``` + +`npubs register` is deliberately narrow — it refuses to do anything if any npub already exists: + +```text +Admin npubs already configured (3). Ask your admin to add your npub. + Your npub: npub1... +``` + +So `register` only ever works once per node. After that, `npubs add` is the command, and it requires admin NIP-98 auth. + +!!! warning "Claim the node during install" + Between the app becoming healthy and the first `npubs register`, the node is unclaimed — anybody who reaches the URL can become admin. See [Deploy on Cloudron](deploy-cloudron.md#bootstrap-the-first-admin). + +--- + +## Command reference + +All of these talk to the auth proxy over NIP-98, so the caller must be a registered npub, and the mutating ones require the `admin` role. + +| Command | Role needed | Notes | +|---|---|---| +| `routstrd npubs list` | any registered | Shows role and name for everyone, and marks your own row with `→ you`. | +| `routstrd npubs register` | none, **once** | Only succeeds while the table is empty. | +| `routstrd npubs add ` | admin | Accepts `npub1...` or a 64-char hex pubkey. `--role admin\|user` (default `user`), `--name`. | +| `routstrd npubs update ` | admin | `--role` and/or `--name`. Passing `--name ""` clears the name. | +| `routstrd npubs delete ` | admin | Revokes access. | + +Names are trimmed and capped at 64 characters. + +`routstrd npubs list` output looks like this: + +```text +Npubs (3): +- npub1qqq...4f2 [admin] "Alice" → you +- npub1xxx...9k7 [user] "Bob" +- npub1zzz...3md [user] +``` + +If your own npub is missing from the list, the CLI tells you whom to send it to: + +```text +Your npub is not in the npub list. Ask an admin to add your npub: + npub1yyy...0pl +``` + +## Underlying HTTP API + +The CLI is a thin wrapper over four endpoints on the auth proxy. Useful for scripting or a custom onboarding form. + +| Method | Path | Auth | Body | +|---|---|---|---| +| `GET` | `/npubs` | any registered npub | — | +| `POST` | `/npubs` | none **if the table is empty**, otherwise admin | `{ "npub": "npub1..." }` or `{ "pubkey": "" }`, plus optional `role` and `name` | +| `PATCH` | `/npubs` | admin | `{ "npub": "npub1..." }` plus `role` and/or `name` (`name: null` clears it) | +| `DELETE` | `/npubs/` | admin | — (also accepts `/npubs?npub=...`) | + +Responses: + +- `GET /npubs` returns `{ "npubs": [ { "npub": "...", "name": "...", "role": "admin" } ] }`. +- Adding a pubkey that is **already registered** returns `409` rather than silently succeeding; use `PATCH` to change an existing entry. +- Every mutation is performed by, and recorded against, the requesting admin. + +!!! note "Revocation is immediate" + Roles and removals are read from the database on **every request** with no caching layer. Deleting an npub stops that member's management access on their next request. Their **API keys are a separate matter** — see below. + +--- + +## What revocation does and does not do + +Both halves of offboarding are separate, and only the first is available through the normal member-facing API. + +**1. Management access — revoked by deleting the npub.** + +```bash +routstrd npubs delete npub1abc...xyz +``` + +Roles and rows are read from the database on **every request** with no caching layer, so the next NIP-98 request from that key is rejected immediately. + +**2. Inference keys — a separate, manually managed thing.** + +The Bearer path validates an `sk-...` key by looking it up in the client records and nothing else. Deleting an npub does **not** touch those rows, so an offboarded member's existing API keys keep working for inference until the client itself is deleted. + +Here the proxy's scoping matters: `/clients`, `/clients/add`, and `/clients/delete` are **strictly owner-scoped** — the proxy filters and authorises by the calling npub. Being an `admin` does not grant access to a *colleague's* client records through those endpoints. So the realistic offboarding paths are: + +- **Have the member delete their own clients** before you remove their npub: `routstrd clients list`, then `routstrd clients delete `. +- **Or do it on the node itself.** Inside the container the CLI talks to the daemon directly on loopback, where there is no auth layer and therefore no ownership filter: + +```bash +cloudron exec --app routstr.example.com +routstrd clients list # every client on the node, all owners +routstrd clients delete +``` + +On the node, client IDs appear **with** their owner suffix (`my-laptop-4f2x9k7`) — see [Connecting Clients](clients.md). That suffix is exactly what tells you which member a client belongs to. + +!!! danger "Removing an npub is not the same as revoking access" + If you skip step 2, a departed member's agents continue to consume the team's wallet. Always pair `npubs delete` with deleting their client records. + +--- + +## Next steps + +- [Connecting Clients](clients.md) — get each member's agents talking to the node. +- [Usage and Model Policy](usage-and-policy.md) — see what each person is spending. diff --git a/docs/teams/troubleshooting.md b/docs/teams/troubleshooting.md new file mode 100644 index 00000000..74c51e38 --- /dev/null +++ b/docs/teams/troubleshooting.md @@ -0,0 +1,161 @@ +# Troubleshooting + +Most team-node problems are one of five things: the daemon never came up, the wrong database, a credential problem, a clock or proxy-header problem, or a networking/timeout issue around streaming. + +--- + +## First triage + +Run these in order. Together they answer "is the node up, does it have people, and is the proxy seeing the right database". + +```bash +# 1. Is the public surface alive? +curl -sS https://team.example.com/health + +# 2. Do both processes run, and what did they log on boot? +cloudron logs --app routstr.example.com | tail -50 # or: docker logs routstr-remote + +# 3. Can the proxy see the database, and does it know your people? +cloudron exec --app routstr.example.com +routstrd-auth validate +``` + +Step 3 is the most informative. Its output ends with a line like `✅ DB accessible. 3 npub(s) registered (1 admin, 2 user).` — if that count is wrong, or the path is wrong, you have found your problem. + +On startup the proxy also logs a summary, and warns loudly if the node is unclaimed: + +```text +routstrd-auth proxy listening on http://0.0.0.0:8008 + Upstream: http://localhost:8009 + DB path: /app/data/routstrd/routstr.db + Registered npubs: 0 + Model allowlist: disabled + Warning: no registered npub/pubkey. The first admin can be registered without auth using POST /npubs. +``` + +!!! warning "`Registered npubs: 0` on an established node is an emergency" + Either the database was lost (usually a missing volume mount) or the proxy is pointed at the wrong file. While the table is empty the node is claimable by anyone who reaches it. Fix it before anything else — see [Lost identity or empty npub list](#lost-identity-or-empty-npub-list). + +--- + +## Error reference + +### Startup + +| Symptom | Cause | Fix | +|---|---|---| +| `Timed out waiting for routstrd to become ready.` | The proxy's launcher waited 120 seconds for the database file to exist **and** for `http://localhost:8009/health` to answer. The daemon is unhealthy or crashing. | Read the daemon's log lines above this one. It is a daemon problem, not a proxy problem. | +| `Database not found at /app/data/routstrd/routstr.db. Make sure routstrd has been initialized (routstrd onboard).` | The proxy shares the daemon's database and never creates the schema. It ran before the daemon had ever initialised. | Start the daemon first (`routstrd start`, or let `supervisord` do it — priority 10 before 20). | +| `Invalid admin Nostr pubkey(s): . Use npub or 64-char hex pubkeys.` | A bootstrap admin variable contains something unparseable. | Fix `ROUTSTRD_AUTH_ADMIN_NPUBS` / `_PUBKEYS` / `_BOOTSTRAP_NPUB`. Use `npub1...` or 64-char hex. | +| Container exits with `Illegal instruction` / `SIGILL` | Bun's default x64 build needs AVX/AVX2, which some hosts and VMs do not expose. | The shipped image already uses the **baseline** build for this reason. If you build your own, keep `BUN_TARGET=bun-linux-x64-baseline`. | +| Proxy crash-loops immediately after a config change | Validation fails, so `start` exits non-zero and `supervisord` restarts it. | Run `routstrd-auth validate` to see the message. | + +### Authentication + +| Response | Meaning | Fix | +|---|---|---| +| `401 Missing Authorization header. Use 'Authorization: Bearer sk-...' or 'Authorization: Nostr '.` | No credential sent, and the path is not public. Remember public paths are `GET`/`HEAD` only. | Send a credential, or use a read method on a public path. | +| `401 Invalid API key.` | The key is not in the client records. | Recover it with `routstrd clients add --name ""`, which prints the existing key. | +| `403 API keys cannot access this endpoint. Use NIP-98 auth from a registered npub/pubkey.` | You used an `sk-...` key on a wallet, node-control, or management endpoint. | Use the CLI, which signs with NIP-98 automatically. | +| `403 Admin access required. Only admin npubs can perform this action.` | The npub is registered but holds the `user` role. | An admin promotes them: `routstrd npubs update --role admin`. | +| `403 This endpoint requires a registered npub/pubkey, but none is configured. Register the first admin with 'routstrd npubs register'.` | The npub table is empty. | Bootstrap the first admin. | +| `403 This endpoint requires NIP-98 auth from a registered npub/pubkey.` | The signature was valid but the pubkey is not in the table. | An admin adds it: `routstrd npubs add `. | +| `401 Invalid Authorization format. Expected 'Bearer sk-...' or 'Nostr '.` | The header used a different scheme or a typo'd prefix. | Fix the prefix. | + +A useful diagnostic: the token **type** determines the error. A `403` naming a specific capability means you authenticated successfully and were then refused by policy. A `401` means you did not authenticate at all. + +### NIP-98 signature rejections + +These all return `401` with a precise reason. + +| Message | Cause | Fix | +|---|---|---| +| `NIP-98 URL tag does not match this request.` | The most common one. The `u` tag holds the URL the client signed, and the proxy reconstructs the public URL from `X-Forwarded-Proto` / `X-Forwarded-Host` (falling back to `Host`). Behind TLS termination, a missing forwarded header makes the reconstructed URL `http://...` while the client signed `https://...`. | Configure the reverse proxy to set both headers. See [Deploy with Docker](deploy-docker.md#put-tls-in-front). | +| `NIP-98 event timestamp is outside the allowed window.` | Events must be within **±60 seconds**. Clock skew. | Sync the client's clock (`timedatectl`, NTP). If only one machine fails, it is that machine. | +| `NIP-98 payload tag is required for requests with a body.` \| `NIP-98 payload tag does not match the request body hash.` | The signed SHA-256 does not match the body that arrived — something rewrote the body in transit. | Check for a middleware, WAF, or forward proxy that re-encodes request bodies. | +| `NIP-98 method tag does not match this request.` | The event was signed for a different method (often a `GET` signature reused on a `POST`). | Sign per request; do not reuse events. | +| `Invalid NIP-98 event signature.` \| `Invalid NIP-98 event kind.` | Corrupted token, or not a kind `27235` event. | Regenerate the request with the CLI. | +| `Invalid NIP-98 token encoding.` \| `Invalid NIP-98 event JSON.` | The base64 payload is truncated — common when a long `Authorization` header is split or truncated by a client. | Check for a header-size limit in the client or proxy. | + +### Client-side CLI messages + +| Message | Meaning | Fix | +|---|---|---| +| `The daemon at rejected this account.` then `Register/authorize this npub on the remote daemon first: ` | The node does not recognise your npub. | Send the printed npub to an admin. | +| `No remote node is set up.` | No `daemonUrl` in `~/.routstrd/config.json`. | `routstrd remote https://team.example.com`. | +| `Your npub is not in the npub list. Ask an admin to add your npub:` | `npubs list` worked (so you *are* registered) but shows you as absent — typically a stale local identity after a config reset. | Re-run `routstrd remote` to display your current npub, and confirm with an admin which one is registered. | +| `Daemon is not running` | The local daemon is unreachable. | Only relevant when running from source; `routstrd start`. | + +--- + +## Common scenarios + +### Streaming responses get cut off mid-answer + +The symptom is a reply that stops abruptly — often after exactly the same number of seconds — with no error from the model. + +This is **not** a node problem. It is response buffering or an idle timeout in a proxy in front of the node. Model turns can be silent for a long time while reasoning or waiting on tools, and intermediaries treat that silence as a dead connection. + +Fix, in order of what actually bites: + +1. **Disable response buffering** where the proxy talks to `8008`. The node already sends `X-Accel-Buffering: no`, but the fronting proxy must also be told (`proxy_buffering off` in nginx). +2. **Raise read timeouts** well above the default (`proxy_read_timeout 3600s`). nginx's 60-second default is the usual culprit. +3. **Do not work around it by exposing the daemon.** Publishing `8009` removes authentication entirely. + +The proxy itself disables Bun's per-request idle timeout (`server.timeout(req, 0)`) precisely because valid streams can be quiet for a long time; anything still timing out is outside the node. + +### Lost identity or empty npub list + +If `routstrd npubs list` is empty after a restart, or `routstrd-auth validate` reports `0 npub(s)`, the node lost its data directory. In Docker this is almost always an **unmounted or changed volume**: `/app/data` is the only persistent path. + +The damage is two-fold: + +- The `nsec` in `routstrd/config.json` is gone, so the node has a **new** identity. +- The npub table is empty, so `POST /npubs` is **unauthenticated again** — anyone who reaches the URL can claim the node. + +Recover in this order: + +1. **Stop the app** so nobody claims it. +2. Restore `/app/data` from a backup, or re-mount the correct volume and restart. +3. If data is unrecoverable, accept the new identity and re-register the first admin — then have everyone whose npub was lost send theirs again, and recreate their clients (keys do not survive either). + +### Model requests return `403` unexpectedly + +Either the sender is genuinely outside the allowlist, or the allowlist is enforcing a **stale** list. + +```bash +cloudron exec --app routstr.example.com +sqlite3 /app/data/routstrd/routstr.db \ + "SELECT substr(value,1,120) FROM sdk_storage WHERE key = 'routstr21Models';" +``` + +- **No output:** the daemon never bootstrapped the list, so the proxy is failing **open** and allowing everything. The `403` is coming from somewhere else. +- **Output present but missing the model you want:** the list is stale. Refresh it with `routstrd clients --manual-refresh`, and check the scheduled job has not been disabled with `--disable-automatic-refresh`. + +Remember the check applies to `POST`/`PUT`/`PATCH` bodies containing a `model` field only, and comparisons are case-sensitive. + +### A member cannot authenticate at all + +Work down this list: + +1. **Are they registered?** `routstrd npubs list` as an admin. +2. **Is their npub the one you registered?** They run `routstrd remote` with no arguments to print the identity actually in their config. A machine with an old or regenerated `nsec` presents a different npub. +3. **Is their clock correct?** Off by more than 60 seconds means every NIP-98 event is rejected. +4. **Are forwarded headers set?** If *everyone* is failing with a URL tag mismatch, it is the reverse proxy, not the people. +5. **Do they have admin-requiring needs?** A `403 Admin access required` is a role problem, not an authentication problem. + +### The node works but nobody can see usage + +Usage is attributed per client, and `/usage` is scoped to the caller's npub. A member with no clients has no usage to show. Aggregate figures require running the CLI **on the node**, where the daemon is unauthenticated and unfiltered: + +```bash +cloudron exec --app routstr.example.com +routstrd top +``` + +--- + +## Next steps + +- [Security Model](security.md) — the rules behind the `401`s and `403`s. +- [Usage and Model Policy](usage-and-policy.md) — allowlist behaviour and its fail-open case. diff --git a/docs/teams/usage-and-policy.md b/docs/teams/usage-and-policy.md new file mode 100644 index 00000000..61ae7e76 --- /dev/null +++ b/docs/teams/usage-and-policy.md @@ -0,0 +1,144 @@ +# Usage and Model Policy + +Two things a team admin cares about: **who is spending what**, and **which models the team may use**. The first is always tracked; the second is an opt-in policy. + +--- + +## Usage tracking + +Every request is attributed to the **client** that made it, and every client is owned by exactly one member. That chain is what produces per-person reporting. + +### As a member + +```bash +routstrd usage # your own usage summary +routstrd top # interactive TUI, alias for 'monitor' +routstrd balance # wallet balance and status +``` + +### As the operator, on the node + +The most complete view is the TUI, run inside the container: + +```bash +cloudron exec --app routstr.example.com +routstrd top +``` + +The **clients** tab shows individual usage, and because each client ID carries the **last 7 characters of its owner's npub**, you can tell teammates apart at a glance even when several of them named their client `my-laptop`: + +```text +my-laptop-4f2x9k7 1,204 req $3.18 +my-laptop-9k7qp2 881 req $2.07 +ci-runner-3md1zz 412 req $0.94 +``` + +Run from the node, this view is **not** filtered to any one member — that is the difference between the CLI on the box and the same CLI on a laptop. + +### Over HTTP + +| Method | Path | Auth | Scope | +|---|---|---|---| +| `GET` | `/usage` | NIP-98 from a registered npub | The **caller's** usage only | +| `GET` | `/usage/summary` | NIP-98 from a registered npub | The **caller's** usage only | + +The proxy forces the scope: it takes the authenticated npub and sets the `npub` query parameter before forwarding to the daemon, so the daemon returns that member's records. There is no query parameter you can pass to widen it — a member cannot read a colleague's usage through the API. Aggregate visibility requires access to the node itself. + +--- + +## The wallet + +The team shares **one node wallet**. Funding it is a single action rather than N personal top-ups, and that is one of the main reasons to run a team node. + +| Endpoint | Required role | +|---|---| +| `/wallet/status`, `/wallet/balance`, `/wallet/mints`, `/wallet/mints/info` | any registered npub | +| `/wallet/receive/cashu`, `/wallet/receive/bolt11` | any registered npub | +| `/wallet/send/cashu`, `/wallet/send/bolt11` | **admin only** | + +Reading the balance and receiving funds are open to every registered member; **moving funds out is admin-only.** Note the consequence: members can see the team's total balance but cannot withdraw from it. API keys can do none of this. + +Funding, mint management, and payment semantics are the daemon's domain and are documented in the provider and client guides rather than repeated here. + +!!! warning "Usage is attributed, not isolated" + There is no per-member balance and no spending cap on this layer. A member's clients spend from the same wallet everyone else does. If you need hard limits, they are not enforced here — track the usage view and manage access accordingly. + +--- + +## Model allowlist + +The proxy can restrict the team to the **Routstr 21 model list**. This is **disabled by default**; enable it with: + +```bash +ROUTSTRD_AUTH_MODEL_ALLOWLIST=true +``` + +When enabled, a request naming a model outside the list is rejected with `403` **before it reaches the daemon**, so no tokens are spent. When disabled, every model passes through untouched. + +### How it works + +```mermaid +flowchart TD + A["routstrd CLI or agent"] --> B["routstrd-auth
checks auth, then model"] + B -->|"allowed"| C["routstrd daemon"] + B -->|"403 not allowed"| A + C --> D["upstream provider"] + E["Nostr kind 38423
Routstr 21 list"] -->|fetched by daemon| F[("sdk_storage
routstr21Models")] + F -->|read on every request| B +``` + +1. The **daemon's** SDK fetches the Routstr 21 list from Nostr (kind `38423` events) and stores it in the shared SQLite database under the `sdk_storage` table, key `routstr21Models`. +2. The **proxy** reads that key from the same database. It has **zero Nostr dependency** — no relay connections, no keys of its own for this purpose. +3. The value is read fresh on every request, with no caching, so a list update takes effect immediately. + +Because the proxy shares the daemon's database, this costs a single indexed key lookup — typically under a millisecond. + +### What is and is not checked + +| Request | Checked | +|---|---| +| `POST` / `PUT` / `PATCH` with a JSON body containing `model` | yes | +| `GET` requests, including `/models` and `/v1/models` | no — public paths, forwarded immediately | +| Management endpoints (`/npubs`, `/clients`, `/usage`) | no — routed to their own handlers before this check | +| Non-JSON bodies | no — no `model` field can be extracted | +| Requests with no `model` field | no — forwarded, and the upstream produces the error | + +Model IDs are compared **case-sensitively**, matching how they are stored in the list. + +### Fail-open behaviour + +If `routstr21Models` is absent — typically because the daemon has not bootstrapped it yet — the proxy **fails open and allows every model**. This is deliberate: the alternative is that a fresh node blocks all traffic until Nostr bootstrapping completes. The trade-off is that an allowlist can be silently ineffective early in a node's life, so verify it after enabling: + +```bash +# confirm the daemon has populated the list +cloudron exec --app routstr.example.com +sqlite3 /app/data/routstrd/routstr.db \ + "SELECT substr(value,1,120) FROM sdk_storage WHERE key = 'routstr21Models';" +``` + +If that returns nothing, the allowlist is not yet meaningful. + +### Performance note + +When enforcement is enabled, the proxy buffers `POST`/`PUT`/`PATCH` request bodies so it can inspect the `model` field. That is a small latency cost on request upload. **Response streaming is unaffected** — SSE and LLM token streams pass through unbuffered, and the proxy explicitly tells intermediaries not to buffer. When the allowlist is disabled, the body is not buffered at all. + +--- + +## Keeping the model list current + +The list is updated by the daemon, not the proxy: + +```bash +routstrd clients --manual-refresh # refresh now +routstrd clients --disable-automatic-refresh # stop the scheduled job +routstrd clients --enable-automatic-refresh # resume it +``` + +If the scheduled refresh job is disabled and nobody runs a manual refresh, the allowlist enforces a **stale** list — and would eventually stop matching newly approved models. + +--- + +## Next steps + +- [Security Model](security.md) — the full endpoint-by-endpoint auth matrix. +- [Troubleshooting](troubleshooting.md) — including "why am I getting a 403". diff --git a/docs/tinfoil-direct-integration.md b/docs/tinfoil-direct-integration.md index c647699a..171adf43 100644 --- a/docs/tinfoil-direct-integration.md +++ b/docs/tinfoil-direct-integration.md @@ -197,9 +197,9 @@ Properties: This is the only architecture that preserves end-to-end encryption from the user to the PPQ/Tinfoil enclave while still letting Routstr mediate payment. The key requirement is that usage/cost metadata must be returned outside the encrypted body, ideally as a response header available before body streaming begins. -## Current Routstr problem +## Original Routstr problem -The current EHBP implementation charges successful EHBP requests at `max_cost_for_model` because Routstr cannot decrypt the response body: +The original EHBP implementation charged successful EHBP requests at `max_cost_for_model` because Routstr could not decrypt the response body: ```text successful EHBP request -> charge full reserved max cost @@ -260,7 +260,13 @@ During local PPQ testing, PPQ responses included this CORS exposure header: Access-Control-Expose-Headers: Ehbp-Response-Nonce, X-Private-Usage-Metrics, X-Encrypted-Usage-Metrics, X-Tinfoil-Usage-Metrics ``` -However, the actual tested non-streaming response did not include any of these usage headers, even when `X-Tinfoil-Request-Usage-Metrics: true` was sent. +However, when this was tested against PPQ's `/private/` endpoint +(`private/gpt-oss-120b`) on 2026-06-21, the non-streaming response did not +include any of these usage headers, even when +`X-Tinfoil-Request-Usage-Metrics: true` was sent. That observation does *not* +hold for the direct Tinfoil enclave upstream that Routstr ships: see +[Usage metrics header format](#usage-metrics-header-format) below, where the +response header and the streaming trailer are both verified present. The decrypted body did include normal OpenAI usage, but only the decrypting Tinfoil client can see that body. @@ -369,7 +375,7 @@ Possible approaches: - PPQ private models are billed per actual input/output tokens. - Private model rates are available from `GET /v1/models?type=all`. -- Current Routstr EHBP billing at max cost is wrong for PPQ private models. +- Max-cost EHBP fallback is wrong for PPQ private models; current code releases/refunds when trusted usage metadata is absent. - Direct Tinfoil integration inside Routstr would enable exact usage billing but would make Routstr see plaintext. - A blind EHBP relay preserves privacy but requires PPQ/Tinfoil to expose usage/cost in plaintext headers/trailers. - The preferred solution is to keep Routstr blind and have PPQ return billing metadata outside the encrypted body. @@ -385,7 +391,9 @@ and `routstr/upstream/ehbp.py`. - Base URL: `https://inference.tinfoil.sh` - Fetches models from the public `GET /v1/models` endpoint (no auth needed). - Parses Tinfoil's pricing (`inputTokenPricePer1M`, `outputTokenPricePer1M`, - `requestPrice`) into the standard `Model`/`Pricing` schema. + `cachedInputTokenPricePer1M`, `requestPrice`) into the standard + `Model`/`Pricing` schema. Cached reads use the cached rate when present, + otherwise the full input rate; cache writes always use the full input rate. - `supports_ehbp = True` — acts as a blind EHBP relay. - `get_ehbp_forwarding_target()` returns a target that includes `X-Tinfoil-Request-Usage-Metrics: true`. @@ -396,9 +404,12 @@ and `routstr/upstream/ehbp.py`. - `routstr/upstream/ehbp.py`: - `parse_tinfoil_usage_metrics()` parses - `prompt=N,completion=N[,total=N][,model=]` into an OpenAI-style - usage dict. The `model` field (added in tinfoilsh/confidential-model-router - PR #385) is extracted as a string. + `prompt=N,completion=N[,total=N][,cached_prompt_tokens=N, + uncached_prompt_tokens=N][,model=][,cost_usd=]` into an + OpenAI-style usage dict. Cache reads map to ``cache_read_input_tokens`` + so ``calculate_cost`` can bill them at the cached rate. The ``model`` + field (added in tinfoilsh/confidential-model-router PR #385) is extracted + as a string; ``cost_usd`` is parsed for logging only. - `_resolve_ehbp_target_url()` overrides the forwarding URL with `X-Tinfoil-Enclave-Url` when the SDK sends it. - `_strip_proxy_headers()` removes `X-Routstr-Model`, @@ -410,8 +421,9 @@ and `routstr/upstream/ehbp.py`. actual served model's pricing is used for cost calculation. - `forward_ehbp_request()` (bearer auth): if `X-Tinfoil-Usage-Metrics` is present in the response header, finalizes with `adjust_payment_for_tokens()` - for exact billing; otherwise falls back to max-cost. Billing uses the - actual served model when it differs from the requested one. + for exact billing; otherwise releases the reservation. The encrypted body + cannot be estimated locally, and the authorization ceiling is not billed. + Billing uses the actual served model when it differs from the requested one. - `forward_ehbp_x_cashu_request()`: if usage is available, computes the refund from actual cost instead of max cost, using the actual served model's pricing when applicable. @@ -425,10 +437,10 @@ and `routstr/upstream/ehbp.py`. |---|---|---| | Bearer, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Exact token cost via `adjust_payment_for_tokens` | | Bearer, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Exact token cost (h11 captures trailers) | -| Bearer, no usage header/trailer | N/A | Max-cost fallback | +| Bearer, no usage header/trailer | N/A | Release reservation; zero charge | | X-Cashu, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Refund = `redeemed - actual_cost` | | X-Cashu, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Refund = `redeemed - actual_cost` (h11 captures trailers) | -| X-Cashu, no usage header/trailer | N/A | Refund = `redeemed - max_cost` | +| X-Cashu, no usage header/trailer | N/A | Full refund | ### Cost response headers @@ -438,13 +450,20 @@ Routstr returns cost info as response headers: | Header | Auth | Description | |---|---|---| -| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Total msats charged for this request | -| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the charge | +| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Settled msats debited for this request | +| `X-Routstr-Computed-Cost-Msats` | Bearer, X-Cashu | Computed usage cost, emitted when it differs from the settled debit | +| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the computed usage | | `X-Routstr-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens | | `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens | +| `X-Routstr-Cache-Read-Msats` | Bearer, X-Cashu | msats attributed to cached input (cache reads) | +| `X-Routstr-Cache-Creation-Msats` | Bearer, X-Cashu | msats attributed to cache creation (0 for Tinfoil today) | The client/Tinfoil SDK can read these headers from the HTTP response without -needing to decrypt the body. +needing to decrypt the body. A duplicate or rejected finalization can therefore +report a zero settled debit while preserving the non-zero computed usage cost. +Normal JSON responses use the same contract: `total_msats` and `charged_msats` +are settled values, while `computed_msats` retains the usage calculation when +it differs. ### Setup @@ -461,9 +480,23 @@ Tinfoil returns usage metrics in the `X-Tinfoil-Usage-Metrics` response header true` is sent. As of tinfoilsh/confidential-model-router PR #385, the format is: ``` -prompt=,completion=,total=,model= +prompt=,completion=,total=[,cached_prompt_tokens=,uncached_prompt_tokens=][,model=][,cost_usd=] ``` +`prompt` is the inclusive prompt total; `cached_prompt_tokens` is the portion +already in Tinfoil's prefix cache and is billed at the model's +`cachedInputTokenPricePer1M` rate (or the full input rate when the model has +no cached rate). `cost_usd` is Tinfoil's own computed request cost and is +currently parsed for observability only — Routstr bills from token counts. + +Note that the header/trailer value is not always a single occurrence: for +streaming responses the trailer is emitted twice, so a client that reads the +trailer directly may see the same `prompt=...,completion=...,...` string twice +in one field, comma-joined. Parsers must be tolerant of the duplicate rather +than assuming exactly one occurrence. `parse_tinfoil_usage_metrics()` is +unaffected: it assigns each `key=value` part as it walks the comma-separated +value, and both occurrences carry identical numbers. + The `model` field carries the actual model name served by the enclave. Routstr uses this to: @@ -490,4 +523,6 @@ back to the requested model's pricing. finalizers for bearer and X-Cashu requests. This provides actual-cost billing today, at the cost of full time-to-last-byte latency for streaming responses. - Whether Tinfoil's `/v1/responses` endpoint also returns usage metrics - headers or trailers. + headers or trailers. Verified: yes — `/v1/responses` returns + `X-Tinfoil-Usage-Metrics` as a plaintext response header, with the same field + set as `/v1/chat/completions` (including `cost_usd`). diff --git a/examples/create_child_keys.py b/examples/create_child_keys.py deleted file mode 100644 index 24f4556d..00000000 --- a/examples/create_child_keys.py +++ /dev/null @@ -1,45 +0,0 @@ -import json -import sys - -import httpx - - -def create_child_keys(base_url: str, api_key: str, count: int = 3) -> list[str]: - headers = {"Authorization": f"Bearer {api_key}"} - - print(f"Requesting {count} child keys from {base_url}...") - - child_keys = [] - - for i in range(count): - try: - response = httpx.post(f"{base_url}/v1/balance/child-key", headers=headers) - if response.status_code == 200: - data = response.json() - child_keys.append(data["api_key"]) - print( - f" [{i + 1}] Created: {data['api_key']} (Cost: {data['cost_msats']} msats)" - ) - else: - print(f" [{i + 1}] Failed: {response.status_code} - {response.text}") - except Exception as e: - print(f" [{i + 1}] Error: {str(e)}") - - return child_keys - - -if __name__ == "__main__": - if len(sys.argv) < 2: - print("Usage: python create_child_keys.py [base_url]") - sys.exit(1) - - auth_key = sys.argv[1] - base_url = sys.argv[2] if len(sys.argv) > 2 else "http://localhost:8000" - - keys = create_child_keys(base_url, auth_key) - - if keys: - print("\nSuccessfully created child keys:") - print(json.dumps(keys, indent=2)) - else: - print("\nNo child keys were created.") diff --git a/examples/tor.py b/examples/tor.py index 7e67bb87..a534c005 100644 --- a/examples/tor.py +++ b/examples/tor.py @@ -7,7 +7,7 @@ from openai import OpenAI client = OpenAI( api_key=os.environ.get("TOKEN"), base_url=os.environ.get("ONION_URL", "http://roustrjfsdgfiueghsklchg.onion/v1"), - http_client=httpx.Client(proxies="socks5://localhost:9050"), + http_client=httpx.Client(proxy="socks5://localhost:9050"), ) print( diff --git a/migrations/versions/3a0fbd387f10_add_refunds_table.py b/migrations/versions/3a0fbd387f10_add_refunds_table.py new file mode 100644 index 00000000..e53e4484 --- /dev/null +++ b/migrations/versions/3a0fbd387f10_add_refunds_table.py @@ -0,0 +1,58 @@ +"""add refunds table + +Revision ID: 3a0fbd387f10 +Revises: a3f1b6c204de +Create Date: 2026-09-16 + +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +revision = "3a0fbd387f10" +down_revision = "a3f1b6c204de" +branch_labels = None +depends_on = None + +OPEN_STATUSES = "status IN ('pending', 'ambiguous')" + + +def upgrade() -> None: + op.create_table( + "refunds", + sa.Column("id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column( + "api_key_hashed_key", sqlmodel.sql.sqltypes.AutoString(), nullable=False + ), + sa.Column("method", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("destination", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("amount_msats", sa.Integer(), nullable=False), + sa.Column("unit", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("mint_url", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("status", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("quote_id", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("token", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("claimed_at", sa.Integer(), nullable=True), + sa.Column("created_at", sa.Integer(), nullable=False), + sa.Column("updated_at", sa.Integer(), nullable=False), + sa.ForeignKeyConstraint(["api_key_hashed_key"], ["api_keys.hashed_key"]), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("ix_refunds_api_key_hashed_key", "refunds", ["api_key_hashed_key"]) + op.create_index("ix_refunds_status", "refunds", ["status"]) + op.create_index( + "ux_refunds_open_per_key", + "refunds", + ["api_key_hashed_key"], + unique=True, + sqlite_where=sa.text(OPEN_STATUSES), + postgresql_where=sa.text(OPEN_STATUSES), + ) + + +def downgrade() -> None: + op.drop_index("ux_refunds_open_per_key", table_name="refunds") + op.drop_index("ix_refunds_status", table_name="refunds") + op.drop_index("ix_refunds_api_key_hashed_key", table_name="refunds") + op.drop_table("refunds") diff --git a/migrations/versions/a3f1b6c204de_add_model_metadata_to_model_paths.py b/migrations/versions/a3f1b6c204de_add_model_metadata_to_model_paths.py new file mode 100644 index 00000000..68baafb2 --- /dev/null +++ b/migrations/versions/a3f1b6c204de_add_model_metadata_to_model_paths.py @@ -0,0 +1,27 @@ +"""add model_metadata to model_paths + +Revision ID: a3f1b6c204de +Revises: e5a6b7c8d9f0 +Create Date: 2026-09-15 21:50:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "a3f1b6c204de" +down_revision = "e5a6b7c8d9f0" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column( + "model_paths", + sa.Column("model_metadata", sa.Text(), nullable=False, server_default="{}"), + ) + + +def downgrade() -> None: + op.drop_column("model_paths", "model_metadata") diff --git a/migrations/versions/e5a6b7c8d9f0_remove_child_keys_and_balance_limits.py b/migrations/versions/e5a6b7c8d9f0_remove_child_keys_and_balance_limits.py new file mode 100644 index 00000000..15298de2 --- /dev/null +++ b/migrations/versions/e5a6b7c8d9f0_remove_child_keys_and_balance_limits.py @@ -0,0 +1,76 @@ +"""Remove child keys and balance limits. + +Removes the child-key feature (parent_key_hash) and the balance-limit +machinery (balance_limit, balance_limit_reset, balance_limit_reset_date) +from api_keys, plus the balance_limit/balance_limit_reset pass-through on +lightning_invoices. + +Data preservation: before dropping the columns, every child key is +converted into a standalone key by clearing parent_key_hash. Child keys +never hold their own balance (they always spent from their parent), so no +funds are lost: the parent keeps its full balance, and the former child +rows are preserved with their total_spent/total_requests history intact. +""" + +import sqlalchemy as sa +from alembic import op + +revision = "e5a6b7c8d9f0" +down_revision = "b4f7a1c9d2e3" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # Convert child keys into standalone keys before dropping the link. + # Their balance is always 0 (they spent from the parent), so this + # cannot strand any funds. + op.execute("UPDATE api_keys SET parent_key_hash = NULL") + + with op.batch_alter_table("api_keys") as batch_op: + batch_op.drop_index("ix_api_keys_parent_key_hash") + batch_op.drop_column("parent_key_hash") + batch_op.drop_column("balance_limit") + batch_op.drop_column("balance_limit_reset") + batch_op.drop_column("balance_limit_reset_date") + + with op.batch_alter_table("lightning_invoices") as batch_op: + batch_op.drop_column("balance_limit") + batch_op.drop_column("balance_limit_reset") + + +def downgrade() -> None: + with op.batch_alter_table("lightning_invoices") as batch_op: + batch_op.add_column(sa.Column("balance_limit", sa.Integer(), nullable=True)) + batch_op.add_column( + sa.Column("balance_limit_reset", sa.String(), nullable=True) + ) + + with op.batch_alter_table("api_keys") as batch_op: + batch_op.add_column( + sa.Column("balance_limit_reset_date", sa.Integer(), nullable=True) + ) + batch_op.add_column( + sa.Column( + "balance_limit_reset", + sa.String(), + nullable=True, + ) + ) + batch_op.add_column(sa.Column("balance_limit", sa.Integer(), nullable=True)) + batch_op.add_column( + sa.Column( + "parent_key_hash", + sa.String(), + nullable=True, + ) + ) + batch_op.create_foreign_key( + "fk_api_keys_parent_key_hash", + "api_keys", + ["parent_key_hash"], + ["hashed_key"], + ) + batch_op.create_index( + "ix_api_keys_parent_key_hash", ["parent_key_hash"], unique=False + ) diff --git a/mkdocs.yml b/mkdocs.yml index 9221c714..0e864f24 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -91,6 +91,15 @@ nav: - Advanced Pricing: provider/advanced-pricing.md - Discovery: provider/discovery.md - Tor Support: provider/tor.md + - Teams (Remote): + - Overview: teams/index.md + - Deploy on Cloudron: teams/deploy-cloudron.md + - Deploy with Docker: teams/deploy-docker.md + - Team Members: teams/team-members.md + - Connecting Clients: teams/clients.md + - Usage and Model Policy: teams/usage-and-policy.md + - Security Model: teams/security.md + - Troubleshooting: teams/troubleshooting.md - API Reference: - Overview: api/overview.md - Authentication: api/authentication.md diff --git a/pyproject.toml b/pyproject.toml index e9abb3cc..9f0eabd1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,27 +1,27 @@ [project] name = "routstr" -version = "0.4.5" +version = "0.4.7" description = "Payment proxy for your LLM endpoint using cashu and nostr." readme = "README.md" requires-python = ">=3.11" dependencies = [ - "fastapi[standard]>=0.115", + "fastapi[standard-no-fastapi-cloud-cli]>=0.141", "aiosqlite>=0.20", - "sqlmodel>=0.0.24", - "httpx[socks]>=0.25.2", - "h11>=0.14", + "sqlmodel>=0.0.42", # Python 3.14 deferred-annotation support + "httpx[socks]>=0.28.1", + "h11>=0.16", "greenlet>=3.2.1", "alembic>=1.13", "python-json-logger>=2.0.0", "cashu>=0.20", "marshmallow>=3.13,<4.0", "websockets>=12.0", - "nostr>=0.0.2", + "nostr-sdk>=0.45.1,<0.46", "mdurl==0.1.2", "pillow>=10", "openai>=1.98.0", - "litellm>=1.55.0", + "litellm>=1.93.0,<1.94", # 1.93 is the first line supporting Python 3.14 ] [dependency-groups] @@ -87,3 +87,32 @@ disallow_untyped_decorators = true [tool.uv.sources] routstr = { workspace = true } + +# Security floors. cashu 0.20.x caps httpx, h11, fastapi, cryptography, +# setuptools and wheel below their patched versions. routstr uses only cashu's +# wallet modules, not its server paths, so the caps are safe to lift here. +[tool.uv] +# coincurve 20's build config uses cmake.verbose, removed in +# scikit-build-core 0.10. This keeps its source build working on Python 3.14. +build-constraint-dependencies = ["scikit-build-core<0.10"] +override-dependencies = [ + # httpx 0.28 (required by litellm 1.93) removed the `proxies` kwarg cashu + # passes on every mint call; routstr/cashu_compat.py restores it for cashu. + "httpx[socks]>=0.28.1,<1.0", + "importlib-metadata>=8.0.0,<9.0", + "h11>=0.16.0", + "fastapi[standard-no-fastapi-cloud-cli]>=0.141", + "cryptography>=49.0.0", + "setuptools>=83.0.0", + "wheel>=0.46.2", +] + +# Transitive deps whose dependents allow the patched version but don't require +# it. Constraints raise the floor without bypassing any upstream pin. +constraint-dependencies = [ + "starlette>=1.3.1", + "httpcore>=1.0.9", # 1.0.8 caps h11<0.15 + # 1.76 is the first grpcio-tools release with CPython 3.14 wheels. + "grpcio>=1.76.0,<2.0.0", + "grpcio-tools>=1.76.0,<2.0.0", +] diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 084972a8..6a896665 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -59,6 +59,23 @@ def calculate_model_cost_score(model: "Model") -> float: return total_cost +def calculate_model_reservation_score(model: "Model") -> float: + """Context-based reservation ceiling for ranking same-model candidates. + + The balance gate reserves on ``sats_pricing.max_cost``, which scales with + ``context_length``. Ranking on the same ceiling keeps the advertised, + routed, and reserved candidate consistent. Lower is better. + """ + from .payment.models import _calculate_usd_max_costs + + try: + _, _, max_cost = _calculate_usd_max_costs(model) + return float(max_cost) + except Exception: + # Pricing shape missing/invalid; per-token score keeps order deterministic + return calculate_model_cost_score(model) + + def get_provider_penalty(provider: "BaseUpstreamProvider") -> float: """Calculate a penalty multiplier for certain providers. @@ -118,9 +135,21 @@ def create_model_mappings( Returns: Tuple of (model_instances, provider_map, unique_models) """ - from .payment.models import _row_to_model + from .payment.models import _row_to_model, has_usable_pricing from .upstream.helpers import resolve_model_alias + def _unusable_price(model: "Model") -> bool: + """A candidate may only route on rates a request can be billed against. + + Mirrors the served-catalog backstop in ``list_models``: a negative or + non-finite rate is not a price, and the cost calculation cannot bill on + one, so every request on the model would be charged the full maximum + reservation instead. Applies to provider-discovered models as well as + persisted overrides — no override row need exist for a malformed price + to be built into the candidate map. + """ + return not has_usable_pricing(model.pricing) + candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {} unique_models: dict[str, "Model"] = {} unique_model_keys: dict[str, str] = {} @@ -208,12 +237,33 @@ def create_model_mappings( # Apply overrides only for this provider's model row. if model_key is not None and model_key in overrides_by_key: override_row, provider_fee = overrides_by_key[model_key] - model_to_use = _row_to_model( - override_row, apply_provider_fee=True, provider_fee=provider_fee - ) + try: + model_to_use = _row_to_model( + override_row, apply_provider_fee=True, provider_fee=provider_fee + ) + except Exception as exc: + # Stored pricing is JSON from whatever wrote the row, so + # converting it can raise. Doing that inside this loop let + # one such row unwind the whole map build: at boot the node + # came up routing nothing, and on a later refresh the map it + # already had went permanently stale. The sibling loop over + # override-only rows already skips and logs such a row. + logger.warning( + "Skipping invalid model override while building model mappings", + extra={ + "model_id": model.id, + "upstream_provider_id": upstream_db_id, + "error": str(exc), + "error_type": type(exc).__name__, + }, + ) + continue else: model_to_use = model + if _unusable_price(model_to_use): + continue + forwarded_model_id = get_effective_forwarded_model_id(model_to_use) # Get all aliases for this model @@ -280,6 +330,8 @@ def create_model_mappings( continue if not model_to_use.enabled: continue + if _unusable_price(model_to_use): + continue forwarded_model_id = get_effective_forwarded_model_id(model_to_use) @@ -323,20 +375,24 @@ def create_model_mappings( def alias_priority(model: "Model", alias: str) -> int: """Rank how strong the mapping of alias->model is. - An exact model ID is authoritative and must be cost-ranked against the - other exact matches before considering forwarded aliases. This keeps a - provider-specific forwarded ID from shadowing a directly available, + A provider that serves the requested ID directly is authoritative and + must be cost-ranked before providers that only reach it through a + forwarded alias, so a forwarded ID cannot shadow a directly available, cheaper model with the requested ID. + + "Directly served" covers both the exact model ID and the same ID behind + a provider prefix (e.g. ``gpt-oss-120b`` on Tinfoil vs + ``openai/gpt-oss-120b`` on OpenRouter). Both name the same model, so + they share the top tier and cost decides between them; otherwise the + provider whose catalog omits the org prefix would always win on ID + spelling regardless of price. """ - if model.id and model.id.lower() == alias: - return 5 + model_base = get_base_model_id(model.id) + if (model.id and model.id.lower() == alias) or model_base.lower() == alias: + return 4 forwarded_model_id = get_effective_forwarded_model_id(model) if forwarded_model_id and forwarded_model_id.lower() == alias: - return 4 - - model_base = get_base_model_id(model.id) - if model_base == alias: return 3 if model.canonical_slug: canonical_base = get_base_model_id(model.canonical_slug) @@ -345,15 +401,19 @@ def create_model_mappings( return 1 for alias, items in candidates.items(): - # Sort key: (priority DESC, cost ASC) - # Using negative cost for DESC sort overall to keep high priority first - def sort_key(item: tuple["Model", "BaseUpstreamProvider"]) -> tuple[int, float]: + # Sort key: (priority DESC, reservation ASC, cost ASC) + # Using negative costs for DESC sort overall to keep high priority first + def sort_key( + item: tuple["Model", "BaseUpstreamProvider"], + ) -> tuple[int, float, float]: model, provider = item priority = alias_priority(model, alias) - cost = calculate_model_cost_score(model) penalty = get_provider_penalty(provider) - adjusted_cost = cost * penalty - return (priority, -adjusted_cost) + # Rank on the reservation ceiling the balance gate enforces, with + # per-token typical-usage cost as tiebreaker + adjusted_reservation = calculate_model_reservation_score(model) * penalty + adjusted_cost = calculate_model_cost_score(model) * penalty + return (priority, -adjusted_reservation, -adjusted_cost) items.sort(key=sort_key, reverse=True) diff --git a/routstr/auth.py b/routstr/auth.py index 588fc2d1..bf29e86e 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -1,12 +1,11 @@ import asyncio import hashlib import math -import random import time import uuid +from contextlib import suppress from contextvars import ContextVar from dataclasses import dataclass -from datetime import datetime from typing import TYPE_CHECKING, Optional from fastapi import HTTPException @@ -29,10 +28,16 @@ from .payment.cost_calculation import ( MaxCostData, calculate_cost, ) +from .redemption_cache import ( + TERMINAL_REDEMPTION_CODES, + CachedRedemptionFailure, + redemption_negative_cache, +) from .wallet import ( classify_redemption_error, credit_balance, deserialize_token_from_string, + resolve_trusted_source_mint, wallet_operation_guard, ) @@ -93,48 +98,6 @@ def _clear_current_reservation(snapshot: ReservationSnapshot) -> None: # PREPAID_BALANCE = int(os.environ.get("PREPAID_BALANCE", "0")) * 1000 # Convert to msats -async def check_and_reset_limit(key: ApiKey, session: AsyncSession) -> bool: - """Checks if a key's balance limit should be reset based on its policy.""" - if key.balance_limit is not None and key.balance_limit_reset: - now = int(time.time()) - reset_date = key.balance_limit_reset_date or 0 - should_reset = False - - if key.balance_limit_reset == "daily": - if ( - datetime.fromtimestamp(now).date() - > datetime.fromtimestamp(reset_date).date() - ): - should_reset = True - elif key.balance_limit_reset == "weekly": - if ( - datetime.fromtimestamp(now).isocalendar()[:2] - > datetime.fromtimestamp(reset_date).isocalendar()[:2] - ): - should_reset = True - elif key.balance_limit_reset == "monthly": - dt_now = datetime.fromtimestamp(now) - dt_reset = datetime.fromtimestamp(reset_date) - if dt_now.year > dt_reset.year or dt_now.month > dt_reset.month: - should_reset = True - - if should_reset: - logger.info( - "Resetting balance limit for key", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "policy": key.balance_limit_reset, - "old_spent": key.total_spent, - }, - ) - key.total_spent = 0 - key.balance_limit_reset_date = now - session.add(key) - await session.flush() - return True - return False - - def redemption_error_to_http_exception(error: Exception) -> HTTPException: """Map a Cashu token redemption failure to a sanitized client-facing error. @@ -166,6 +129,50 @@ def redemption_error_to_http_exception(error: Exception) -> HTTPException: ) +def _cached_failure_to_http_exception( + failure: CachedRedemptionFailure, +) -> HTTPException: + """Rebuild the exact error envelope the original mint-backed failure produced.""" + return HTTPException( + status_code=failure.status_code, + detail={ + "error": { + "message": failure.message, + "type": failure.error_type, + "code": failure.code, + } + }, + ) + + +def _maybe_cache_terminal_redemption_failure(hashed_key: str, error: Exception) -> None: + """Record a redemption failure in the negative cache if it can never succeed. + + Transient classifications (mint unreachable, rate-limited) are never + cached — only codes in TERMINAL_REDEMPTION_CODES, which are permanent + properties of the token itself. + """ + classified = classify_redemption_error(error) + if classified is None: + return + error_type, status_code, message, code = classified + if code not in TERMINAL_REDEMPTION_CODES: + return + redemption_negative_cache.put( + hashed_key, + CachedRedemptionFailure( + status_code=status_code, + error_type=error_type, + message=message, + code=code, + ), + ) + logger.info( + "Cached terminal redemption failure; further attempts rejected locally", + extra={"key_hash": hashed_key[:8] + "...", "code": code}, + ) + + async def validate_bearer_key( bearer_key: str, session: AsyncSession, @@ -200,7 +207,7 @@ async def _validate_bearer_key_locked( Validates the provided API key using SQLModel. If it's a cashu key, it redeems it and stores its hash and balance. Otherwise checks if the hash of the key exists. - Includes a balance check against min_cost for limited keys. + Checks the key's available balance against min_cost when required. """ logger.debug( "Starting bearer key validation", @@ -265,42 +272,19 @@ async def _validate_bearer_key_locked( }, ) - # Check and reset limit if needed - await check_and_reset_limit(existing_key, session) - - # Early check: Billing balance check (Parent balance) - billing_key = await get_billing_key(existing_key, session) - if min_cost > 0 and billing_key.total_balance < min_cost: + # Early check: Billing balance check + if min_cost > 0 and existing_key.total_balance < min_cost: logger.warning( "Insufficient billing balance during validation", extra={ "key_hash": existing_key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", - "balance": billing_key.total_balance, + "balance": existing_key.total_balance, "required": min_cost, }, ) raise HTTPException( status_code=402, - detail=_model_balance_error(min_cost, billing_key.total_balance), - ) - - # Early check: Spending limit check (Child key limit) - if ( - min_cost > 0 - and existing_key.balance_limit is not None - and existing_key.total_spent + existing_key.reserved_balance + min_cost - > existing_key.balance_limit - ): - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"Balance limit exceeded: {existing_key.balance_limit} mSats limit. {existing_key.total_spent} already spent ({existing_key.reserved_balance} reserved), {min_cost} minimum required for this model.", - "type": "insufficient_quota", - "code": "balance_limit_exceeded", - } - }, + detail=_model_balance_error(min_cost, existing_key.total_balance), ) return existing_key @@ -379,6 +363,16 @@ async def _validate_bearer_key_locked( return existing_key + if cached_failure := redemption_negative_cache.get(hashed_key): + logger.info( + "Rejecting known-dead Cashu token from negative cache", + extra={ + "key_hash": hashed_key[:8] + "...", + "code": cached_failure.code, + }, + ) + raise _cached_failure_to_http_exception(cached_failure) + logger.info( "Creating new Cashu token entry", extra={ @@ -387,24 +381,20 @@ async def _validate_bearer_key_locked( "has_expiry_time": bool(key_expiry_time), }, ) - if token_obj.mint == settings.primary_mint: - if token_obj.unit != settings.primary_mint_unit: - raise redemption_error_to_http_exception( - ValueError( - "Cashu token unit does not match the configured primary " - f"mint unit: expected {settings.primary_mint_unit}, " - f"got {token_obj.unit}" - ) + token_mint = resolve_trusted_source_mint(token_obj.mint) or token_obj.mint + if ( + token_mint == settings.primary_mint + and token_obj.unit != settings.primary_mint_unit + ): + raise redemption_error_to_http_exception( + ValueError( + "Cashu token unit does not match the configured primary " + f"mint unit: expected {settings.primary_mint_unit}, " + f"got {token_obj.unit}" ) - refund_currency = token_obj.unit - refund_mint_url = settings.primary_mint - elif token_obj.mint in settings.cashu_mints: - refund_currency = token_obj.unit - refund_mint_url = token_obj.mint - else: - # Foreign tokens are swapped into the configured primary mint. - refund_currency = settings.primary_mint_unit - refund_mint_url = settings.primary_mint + ) + refund_currency = token_obj.unit + refund_mint_url = token_mint new_key = ApiKey( hashed_key=hashed_key, @@ -450,14 +440,30 @@ async def _validate_bearer_key_locked( "AUTH: credit_balance returned successfully", extra={"msats": msats} ) except Exception as credit_error: - logger.error( + classification = classify_redemption_error(credit_error) + expected_codes = { + "cashu_token_already_spent", + "cashu_source_mint_unreachable", + "cashu_mint_unreachable", + "cashu_mint_rate_limited", + "cashu_mint_timeout", + } + log = ( + logger.info + if classification is not None + and classification[3] in expected_codes + else logger.error + ) + log( "AUTH: credit_balance failed", extra={ "error": str(credit_error), "error_type": type(credit_error).__name__, + "error_code": classification[3] if classification else None, }, ) await session.rollback() + _maybe_cache_terminal_redemption_failure(hashed_key, credit_error) raise redemption_error_to_http_exception(credit_error) from credit_error if msats <= 0: @@ -544,27 +550,6 @@ async def _validate_bearer_key_locked( ) -async def get_billing_key(key: ApiKey, session: AsyncSession) -> ApiKey: - """Returns the key that should be charged for the request.""" - if key.parent_key_hash: - parent = await session.get(ApiKey, key.parent_key_hash) - if parent: - # We want to keep the total_requests and total_spent on the child key - # but use the balance and reserved_balance of the parent. - # However, pay_for_request updates reserved_balance and total_requests. - # To stay simple, we charge the parent's balance and update parent's total_requests. - return parent - else: - logger.error( - "Parent key not found for child key", - extra={ - "child_key_hash": key.hashed_key[:8] + "...", - "parent_key_hash": key.parent_key_hash[:8] + "...", - }, - ) - return key - - async def pay_for_request( key: ApiKey, cost_per_request: int, session: AsyncSession ) -> int: @@ -572,7 +557,7 @@ async def pay_for_request( # Ensure cost_per_request is at least the minimum allowed request cost cost_per_request = max(cost_per_request, settings.min_request_msat) - billing_key = await get_billing_key(key, session) + billing_key = key logger.info( "Processing payment for request", @@ -631,35 +616,6 @@ async def pay_for_request( }, ) - # Check balance limit for child keys (or any key with a limit) - if key.balance_limit is not None: - await check_and_reset_limit(key, session) - - if ( - key.total_spent + key.reserved_balance + cost_per_request - > key.balance_limit - ): - logger.warning( - "Balance limit exceeded", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "total_spent": key.total_spent, - "reserved": key.reserved_balance, - "balance_limit": key.balance_limit, - "required": cost_per_request, - }, - ) - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": f"Balance limit exceeded: {key.balance_limit} mSats limit. {key.total_spent} already spent ({key.reserved_balance} reserved), {cost_per_request} required for this request.", - "type": "insufficient_quota", - "code": "balance_limit_exceeded", - } - }, - ) - logger.debug( "Charging base cost for request", extra={ @@ -695,13 +651,19 @@ async def pay_for_request( result = await session.exec(stmt) # type: ignore[call-overload] if result.rowcount == 0: - logger.error( - "Concurrent request depleted balance", + await session.refresh(billing_key) + total_balance = billing_key.balance + reserved_balance = billing_key.reserved_balance + available_balance = max(0, total_balance - reserved_balance) + logger.warning( + "Concurrent request depleted available balance", extra={ "key_hash": key.hashed_key[:8] + "...", "billing_key_hash": billing_key.hashed_key[:8] + "...", "required_cost": cost_per_request, - "current_balance": billing_key.balance, + "total_balance": total_balance, + "reserved_balance": reserved_balance, + "available_balance": available_balance, }, ) @@ -709,59 +671,14 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.balance} available.", + "message": f"Insufficient balance: {cost_per_request} mSats required. {available_balance} available.", "type": "insufficient_quota", "code": "insufficient_balance", + "available_balance": available_balance, } }, ) - # Also increment total_requests and reserved_balance on the child key if it's different. - # The balance_limit guard is enforced atomically here — the Python pre-check above - # is a fast-path rejection only and provides no concurrency guarantee. - if billing_key.hashed_key != key.hashed_key: - child_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where( - (col(ApiKey.balance_limit).is_(None)) - | ( - col(ApiKey.total_spent) - + col(ApiKey.reserved_balance) - + cost_per_request - <= col(ApiKey.balance_limit) - ) - ) - .values( - total_requests=col(ApiKey.total_requests) + 1, - reserved_balance=col(ApiKey.reserved_balance) + cost_per_request, - reserved_at=reserved_at_now, - ) - ) - child_result = await session.exec(child_stmt) # type: ignore[call-overload] - - if child_result.rowcount == 0: - # Build the error before rollback expires ORM attributes. - limit_message = ( - f"Balance limit exceeded: {key.balance_limit} mSats limit. " - f"{key.total_spent} already spent ({key.reserved_balance} reserved), " - f"{cost_per_request} required for this request." - ) - # The parent reservation update already ran in this transaction. - # Roll it back before failover code attempts to restore the previous - # reservation; otherwise that later commit can persist both updates. - await session.rollback() - raise HTTPException( - status_code=402, - detail={ - "error": { - "message": limit_message, - "type": "insufficient_quota", - "code": "balance_limit_exceeded", - } - }, - ) - session.add( ReservationRelease( id=reservation.release_id, @@ -807,10 +724,12 @@ async def pay_for_request( _clear_current_reservation(reservation) raise + # The reservation is durable; keep its lease fresh for the whole request + # lifetime (upstream header waits, non-streaming and streaming alike). + _start_reservation_heartbeat(reservation) + try: await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) except Exception: # The reservation transaction is already committed and durable. Logging # refresh failures must not make the caller treat it as unreserved. @@ -882,9 +801,6 @@ async def _validate_reservation_snapshot( persisted_key = await session.get(ApiKey, snapshot.key_hash) if persisted_key is None: raise RuntimeError("Billing reservation key no longer exists") - expected_billing_hash = persisted_key.parent_key_hash or persisted_key.hashed_key - if snapshot.billing_key_hash != expected_billing_hash: - raise RuntimeError("Billing reservation does not belong to this billing key") record = await session.get(ReservationRelease, snapshot.release_id) if ( @@ -897,6 +813,84 @@ async def _validate_reservation_snapshot( raise RuntimeError("Billing reservation record does not match the request") +async def renew_reservation( + snapshot: ReservationSnapshot, session: AsyncSession +) -> bool: + """Push an active reservation's lease forward so the sweeper skips it. + + ``ReservationRelease.created_at`` doubles as the lease timestamp: the + stale-reservation sweeper releases reservations whose ``created_at`` is + older than the timeout, so a long-lived stream must renew it periodically + or lose its reservation mid-flight (and finish uncharged, since release is + terminal). Returns False once the reservation reached a terminal state. + """ + result = await session.exec( # type: ignore[call-overload] + update(ReservationRelease) + .where(col(ReservationRelease.id) == snapshot.release_id) + .where(col(ReservationRelease.status) == "active") + .values(created_at=int(time.time())) + ) + await session.commit() + return bool(result.rowcount == 1) + + +# One heartbeat task per in-flight reservation, keyed by release id. Started +# when the reservation is created and stopped when it reaches a terminal +# state, so every request path — header waits, non-streaming, streaming — is +# covered for its whole lifetime. +_reservation_heartbeats: dict[str, "asyncio.Task[None]"] = {} + + +def _start_reservation_heartbeat(snapshot: ReservationSnapshot) -> None: + """Keep an in-flight reservation's lease fresh until it is finalized. + + Spawns a background task that renews the lease every third of the stale + timeout using its own session, so requests longer than + ``STALE_RESERVATION_TIMEOUT_SECONDS`` are not swept and finish charged. + The task stops on its own once the reservation reaches a terminal state + or its owning request task finishes; terminal transitions also stop it + explicitly. Binding renewal to the owner's lifetime guarantees the sweeper + can always recover a reservation whose request died without finalizing — + a detached heartbeat would otherwise renew it forever and lock the funds. + """ + interval = max(1, settings.stale_reservation_timeout_seconds // 3) + owner = asyncio.current_task() + + async def beat() -> None: + try: + while True: + await asyncio.sleep(interval) + if owner is None or owner.done(): + # Request control is gone; let the lease expire so the + # sweeper can release the reservation if no terminal + # transition ever ran. + return + try: + async with create_session() as session: + if not await renew_reservation(snapshot, session): + return + except Exception: + logger.exception( + "Failed to renew billing reservation lease", + extra={"release_id": snapshot.release_id}, + ) + finally: + _reservation_heartbeats.pop(snapshot.release_id, None) + + _reservation_heartbeats[snapshot.release_id] = asyncio.create_task(beat()) + + +async def _stop_reservation_heartbeat(release_id: str) -> None: + """Cancel and await a reservation's heartbeat so no renewal overlaps + finalization.""" + task = _reservation_heartbeats.pop(release_id, None) + if task is None: + return + task.cancel() + with suppress(asyncio.CancelledError): + await task + + async def get_reservation_snapshot( key: ApiKey, session: AsyncSession ) -> ReservationSnapshot: @@ -908,6 +902,60 @@ async def get_reservation_snapshot( return snapshot +async def _repair_corrupt_reservation( + snapshot: ReservationSnapshot, + session: AsyncSession, + *, + decrement_requests: bool, +) -> bool: + """Terminalize a reservation without subtracting uncertain aggregates.""" + transition = ( + update(ReservationRelease) + .where(col(ReservationRelease.id) == snapshot.release_id) + .where(col(ReservationRelease.status) == "active") + .where(col(ReservationRelease.key_hash) == snapshot.key_hash) + .where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash) + .where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats) + .values(status="released") + ) + result = await session.exec(transition) # type: ignore[call-overload] + if result.rowcount != 1: + await session.rollback() + return False + + if decrement_requests: + for key_hash in {snapshot.billing_key_hash, snapshot.key_hash}: + request_result = await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == key_hash) + .values( + total_requests=case( + ( + col(ApiKey.total_requests) > 0, + col(ApiKey.total_requests) - 1, + ), + else_=0, + ) + ) + ) + if request_result.rowcount != 1: + await session.rollback() + return False + + await session.commit() + logger.error( + "Released corrupt reservation without aggregate subtraction", + extra={ + "reservation_id": snapshot.release_id, + "billing_key_hash": snapshot.billing_key_hash[:8] + "...", + "reserved_msats": snapshot.reserved_msats, + }, + ) + await _stop_reservation_heartbeat(snapshot.release_id) + _clear_current_reservation(snapshot) + return True + + async def _transition_reservation_to_released( snapshot: ReservationSnapshot, session: AsyncSession, @@ -928,7 +976,7 @@ async def _transition_reservation_to_released( if transition_result.rowcount != 1: await session.rollback() existing = await session.get(ReservationRelease, snapshot.release_id) - return bool( + already_released = bool( idempotent_success and existing is not None and existing.status == "released" @@ -936,6 +984,9 @@ async def _transition_reservation_to_released( and existing.billing_key_hash == snapshot.billing_key_hash and existing.reserved_msats == snapshot.reserved_msats ) + if already_released: + await _stop_reservation_heartbeat(snapshot.release_id) + return already_released values: dict[str, object] = { "reserved_balance": col(ApiKey.reserved_balance) - snapshot.reserved_msats, @@ -959,23 +1010,12 @@ async def _transition_reservation_to_released( result = await session.exec(release_stmt) # type: ignore[call-overload] if result.rowcount != 1: await session.rollback() - return False - - if snapshot.billing_key_hash != snapshot.key_hash: - child_release_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == snapshot.key_hash) - .where(col(ApiKey.reserved_balance) >= snapshot.reserved_msats) - .values(**values) + return await _repair_corrupt_reservation( + snapshot, session, decrement_requests=decrement_requests ) - child_result = await session.exec( # type: ignore[call-overload] - child_release_stmt - ) - if child_result.rowcount != 1: - await session.rollback() - return False await session.commit() + await _stop_reservation_heartbeat(snapshot.release_id) _clear_current_reservation(snapshot) return True @@ -1011,6 +1051,9 @@ async def _claim_reservation_for_charge( ) result = await session.exec(statement) # type: ignore[call-overload] if result.rowcount == 1: + # The claim is not committed yet — the heartbeat must keep running + # until the surrounding charge transaction commits, or a rollback + # would restore an active reservation with no lease renewal. _clear_current_reservation(snapshot) return True @@ -1018,6 +1061,49 @@ async def _claim_reservation_for_charge( return False +async def _charge_reservation_rows( + session: AsyncSession, + *, + billing_key_hash: str, + reserved_msats: int, + charge_msats: int, + extra_billing_guards: tuple = (), +) -> bool: + """Release the reserved amount and record the charge on the key + inside the caller's transaction. + + Guarded subtraction replaces defensive clamping: the row must still hold + the full reserved amount, otherwise the whole transaction rolls back and + nothing is charged. A violated invariant must never silently erase the + aggregate reservations of sibling requests. Returns False after rollback. + """ + billing_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == billing_key_hash) + .where(col(ApiKey.reserved_balance) >= reserved_msats) + .values( + reserved_balance=col(ApiKey.reserved_balance) - reserved_msats, + reserved_at=case( + ( + col(ApiKey.reserved_balance) - reserved_msats > 0, + col(ApiKey.reserved_at), + ), + else_=None, + ), + balance=col(ApiKey.balance) - charge_msats, + total_spent=col(ApiKey.total_spent) + charge_msats, + ) + ) + for guard in extra_billing_guards: + billing_stmt = billing_stmt.where(guard) + result = await session.exec(billing_stmt) # type: ignore[call-overload] + if result.rowcount != 1: + await session.rollback() + return False + + return True + + async def adjust_payment_for_tokens( key: ApiKey, response_data: dict, @@ -1039,7 +1125,7 @@ async def adjust_payment_for_tokens( The response's usage object is normalized with the default union parser in ``calculate_cost``. """ - billing_key = await get_billing_key(key, session) + billing_key = key reservation = reservation_snapshot or await get_reservation_snapshot(key, session) await _validate_reservation_snapshot( key, reservation, session, require_active=False @@ -1048,6 +1134,10 @@ async def adjust_payment_for_tokens( # changed the caller's original estimate. deducted_max_cost = reservation.reserved_msats model = response_data.get("model", "unknown") + # Failure paths log after a rollback has expired the ORM instances, so + # capture the identifiers as plain strings up front. + key_log_hash = key.hashed_key[:8] + "..." + billing_log_hash = billing_key.hashed_key[:8] + "..." logger.debug( "Starting payment adjustment for tokens", @@ -1072,8 +1162,8 @@ async def adjust_payment_for_tokens( if released else "Reservation was already finalized; fallback skipped", extra={ - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", + "key_hash": key_log_hash, + "billing_key_hash": billing_log_hash, "deducted_max_cost": deducted_max_cost, }, ) @@ -1082,8 +1172,8 @@ async def adjust_payment_for_tokens( "Failed to release reservation in fallback", extra={ "error": str(e), - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", + "key_hash": key_log_hash, + "billing_key_hash": billing_log_hash, }, ) @@ -1101,12 +1191,27 @@ async def adjust_payment_for_tokens( calculated_cost = await calculate_cost( response_data, deducted_max_cost, model_obj, provider_fee ) - if not isinstance(calculated_cost, CostDataError): - if not await _claim_reservation_for_charge(reservation, session): - # A prior charge or release already owns this reservation. Returning - # the calculated metadata is safe; the aggregate balances must not - # be modified a second time. - return calculated_cost.dict() + if isinstance(calculated_cost, CostDataError): + # Content was already served, so release instead of raising a 400. + logger.error( + "Cost calculation error during payment adjustment, releasing reservation", + extra={ + "key_hash": key_log_hash, + "model": model, + "error_message": calculated_cost.message, + "error_code": calculated_cost.code, + }, + ) + calculated_cost = MaxCostData( + base_msats=0, input_msats=0, output_msats=0, total_msats=0 + ) + + if not await _claim_reservation_for_charge(reservation, session): + # A prior charge or release already owns this reservation. Returning + # the calculated metadata is safe; the aggregate balances must not + # be modified a second time. + calculated_cost.charged_msats = 0 + return calculated_cost.dict() match calculated_cost: case MaxCostData() as cost: @@ -1119,78 +1224,32 @@ async def adjust_payment_for_tokens( "max_cost": cost.total_msats, }, ) - # Finalize by releasing reservation and charging max cost - if billing_key.reserved_balance < deducted_max_cost: - logger.error( - "reserved_balance below deducted_max_cost before MaxCost finalization — clamping to 0", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", - "reserved_balance": billing_key.reserved_balance, - "deducted_max_cost": deducted_max_cost, - "total_cost_msats": cost.total_msats, - "balance": billing_key.balance, - "total_spent": billing_key.total_spent, - "model": model, - }, - ) - - safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost, - ), - else_=0, + # Finalize by releasing the reservation and charging max cost. + charged = await _charge_reservation_rows( + session, + billing_key_hash=billing_key.hashed_key, + reserved_msats=deducted_max_cost, + charge_msats=cost.total_msats, ) - - finalize_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == billing_key.hashed_key) - .values( - reserved_balance=safe_reserved, - balance=col(ApiKey.balance) - cost.total_msats, - total_spent=col(ApiKey.total_spent) + cost.total_msats, - ) - ) - result = await session.exec(finalize_stmt) # type: ignore[call-overload] - - # Also update total_spent and reserved_balance on the child key if it's different - if billing_key.hashed_key != key.hashed_key: - child_safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost, - ), - else_=0, - ) - child_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values( - total_spent=col(ApiKey.total_spent) + cost.total_msats, - reserved_balance=child_safe_reserved, - ) - ) - await session.exec(child_stmt) # type: ignore[call-overload] - - await session.commit() - if result.rowcount == 0: + if charged: + await session.commit() + await _stop_reservation_heartbeat(reservation.release_id) + if not charged: logger.error( "Failed to finalize max-cost payment - retrying reservation release", extra={ - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", + "key_hash": key_log_hash, + "billing_key_hash": billing_log_hash, "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": billing_key.reserved_balance, "total_cost": cost.total_msats, "model": model, }, ) + cost.charged_msats = 0 await release_reservation_only() else: + cost.charged_msats = cost.total_msats await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) logger.info( "Max cost payment finalized", extra={ @@ -1254,63 +1313,30 @@ async def adjust_payment_for_tokens( "model": model, }, ) - if billing_key.reserved_balance < deducted_max_cost: + if not await _charge_reservation_rows( + session, + billing_key_hash=billing_key.hashed_key, + reserved_msats=deducted_max_cost, + charge_msats=total_cost_msats, + ): logger.error( - "reserved_balance below deducted_max_cost on exact-cost finalization — clamping to 0", + "Failed to finalize exact-cost payment - releasing reservation", extra={ - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", - "reserved_balance": billing_key.reserved_balance, + "key_hash": key_log_hash, + "billing_key_hash": billing_log_hash, "deducted_max_cost": deducted_max_cost, - "total_cost_msats": total_cost_msats, - "balance": billing_key.balance, - "total_spent": billing_key.total_spent, + "total_cost": total_cost_msats, "model": model, }, ) - - exact_safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost, - ), - else_=0, - ) - - finalize_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == billing_key.hashed_key) - .values( - reserved_balance=exact_safe_reserved, - balance=col(ApiKey.balance) - total_cost_msats, - total_spent=col(ApiKey.total_spent) + total_cost_msats, - ) - ) - await session.exec(finalize_stmt) # type: ignore[call-overload] - - # Also update total_spent and reserved_balance on the child key if it's different - if billing_key.hashed_key != key.hashed_key: - child_exact_safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost, - ), - else_=0, - ) - child_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values( - total_spent=col(ApiKey.total_spent) + total_cost_msats, - reserved_balance=child_exact_safe_reserved, - ) - ) - await session.exec(child_stmt) # type: ignore[call-overload] + cost.charged_msats = 0 + await release_reservation_only() + return cost.dict() await session.commit() + await _stop_reservation_heartbeat(reservation.release_id) + cost.charged_msats = total_cost_msats await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) await _accumulate_fee(total_cost_msats) payments_logger.info( "FINALIZE", @@ -1333,8 +1359,8 @@ async def adjust_payment_for_tokens( # actual cost exceeded discounted reservation (due to tolerance_percentage) if cost_difference > 0: - # Lock the billing row so the parent and child record the same - # database-determined charge under concurrent finalizations. + # Lock the key row so concurrent finalizations use the same + # database-determined charge. actual_charge_msats = 0 for attempt in range(5): locked_billing_key = ( @@ -1346,50 +1372,67 @@ async def adjust_payment_for_tokens( ) ).one() observed_balance = locked_billing_key.balance - actual_charge_msats = min(observed_balance, total_cost_msats) - overrun_safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost, - ), - else_=0, - ) - finalize_result = await session.exec( # type: ignore[call-overload] - update(ApiKey) - .where(col(ApiKey.hashed_key) == billing_key.hashed_key) - .where(col(ApiKey.balance) == observed_balance) - .values( - reserved_balance=overrun_safe_reserved, - balance=col(ApiKey.balance) - actual_charge_msats, - total_spent=col(ApiKey.total_spent) + actual_charge_msats, + observed_reserved = locked_billing_key.reserved_balance + # An overrun may only spend this request's own reservation + # plus funds no other in-flight request has reserved. + # Charging against the raw balance would consume sibling + # reservations and drive the available balance negative. + if observed_reserved < deducted_max_cost: + # Invariant violated — never clamp and charge anyway, + # that would erase sibling reservations. Release only. + logger.error( + "reserved_balance below reservation on overrun finalization — releasing without charge", + extra={ + "key_hash": key_log_hash, + "billing_key_hash": billing_log_hash, + "reserved_balance": observed_reserved, + "deducted_max_cost": deducted_max_cost, + "total_cost_msats": total_cost_msats, + "model": model, + }, ) - ) - if finalize_result.rowcount == 1: + await session.rollback() + cost.charged_msats = 0 + await release_reservation_only() + return cost.dict() + sibling_reserved = observed_reserved - deducted_max_cost + chargeable_msats = max(0, observed_balance - sibling_reserved) + actual_charge_msats = min(chargeable_msats, total_cost_msats) + if await _charge_reservation_rows( + session, + billing_key_hash=billing_key.hashed_key, + reserved_msats=deducted_max_cost, + charge_msats=actual_charge_msats, + extra_billing_guards=( + col(ApiKey.balance) == observed_balance, + col(ApiKey.reserved_balance) == observed_reserved, + ), + ): break - await session.rollback() if not await _claim_reservation_for_charge(reservation, session): + cost.charged_msats = 0 return cost.dict() else: await session.rollback() raise RuntimeError("Could not atomically finalize cost overrun") - if billing_key.hashed_key != key.hashed_key: - child_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values( - reserved_balance=overrun_safe_reserved, - total_spent=col(ApiKey.total_spent) + actual_charge_msats, - ) - ) - await session.exec(child_stmt) # type: ignore[call-overload] - await session.commit() + await _stop_reservation_heartbeat(reservation.release_id) await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) - cost.total_msats = actual_charge_msats + cost.charged_msats = actual_charge_msats + if actual_charge_msats < total_cost_msats: + logger.warning( + "Cost overrun exceeded chargeable funds — shortfall written off", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "billing_key_hash": billing_key.hashed_key[:8] + "...", + "actual_cost_msats": total_cost_msats, + "charged_msats": actual_charge_msats, + "shortfall_msats": total_cost_msats - actual_charge_msats, + "model": model, + }, + ) logger.info( "Finalized payment with additional charge", extra={ @@ -1432,80 +1475,33 @@ async def adjust_payment_for_tokens( }, ) - if billing_key.reserved_balance < deducted_max_cost: - logger.error( - "reserved_balance below deducted_max_cost on refund finalization — clamping to 0", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", - "reserved_balance": billing_key.reserved_balance, - "deducted_max_cost": deducted_max_cost, - "total_cost_msats": total_cost_msats, - "refund_amount": refund, - "balance": billing_key.balance, - "total_spent": billing_key.total_spent, - "model": model, - }, - ) - - refund_safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost, - ), - else_=0, + charged = await _charge_reservation_rows( + session, + billing_key_hash=billing_key.hashed_key, + reserved_msats=deducted_max_cost, + charge_msats=total_cost_msats, ) + if charged: + await session.commit() + await _stop_reservation_heartbeat(reservation.release_id) - refund_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == billing_key.hashed_key) - .values( - reserved_balance=refund_safe_reserved, - balance=col(ApiKey.balance) - total_cost_msats, - total_spent=col(ApiKey.total_spent) + total_cost_msats, - ) - ) - result = await session.exec(refund_stmt) # type: ignore[call-overload] - - # Also update total_spent and reserved_balance on the child key if it's different - if billing_key.hashed_key != key.hashed_key: - child_refund_safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= deducted_max_cost, - col(ApiKey.reserved_balance) - deducted_max_cost, - ), - else_=0, - ) - child_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values( - total_spent=col(ApiKey.total_spent) + total_cost_msats, - reserved_balance=child_refund_safe_reserved, - ) - ) - await session.exec(child_stmt) # type: ignore[call-overload] - - await session.commit() - - if result.rowcount == 0: + if not charged: logger.error( "Failed to finalize payment - releasing reservation", extra={ - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", + "key_hash": key_log_hash, + "billing_key_hash": billing_log_hash, "deducted_max_cost": deducted_max_cost, - "current_reserved_balance": billing_key.reserved_balance, "total_cost": total_cost_msats, "model": model, }, ) + cost.charged_msats = 0 await release_reservation_only() else: cost.total_msats = total_cost_msats + cost.charged_msats = total_cost_msats await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) logger.info( "Refund processed successfully", @@ -1540,94 +1536,10 @@ async def adjust_payment_for_tokens( return cost.dict() - case CostDataError() as error: - logger.error( - "Cost calculation error during payment adjustment - releasing reservation", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "model": model, - "error_message": error.message, - "error_code": error.code, - }, - ) - await release_reservation_only() - - raise HTTPException( - status_code=400, - detail={ - "error": { - "message": error.message, - "type": "invalid_request_error", - "code": error.code, - } - }, - ) # All calculate_cost variants are handled above. raise AssertionError("Unreachable: unhandled calculate_cost result") -async def periodic_key_reset() -> None: - """Background task to reset key limits based on their policy.""" - from .core.db import create_session - - while True: - try: - interval = 3600 # Run every hour - jitter = 300 - await asyncio.sleep(interval + random.uniform(0, jitter)) - except asyncio.CancelledError: - break - - try: - async with create_session() as session: - # Find all keys that have a reset policy - stmt = select(ApiKey).where(ApiKey.balance_limit_reset.is_not(None)) # type: ignore - keys = (await session.exec(stmt)).all() - - now = int(time.time()) - updated_count = 0 - - for key in keys: - reset_date = key.balance_limit_reset_date or 0 - should_reset = False - - if key.balance_limit_reset == "daily": - if ( - datetime.fromtimestamp(now).date() - > datetime.fromtimestamp(reset_date).date() - ): - should_reset = True - elif key.balance_limit_reset == "weekly": - if ( - datetime.fromtimestamp(now).isocalendar()[:2] - > datetime.fromtimestamp(reset_date).isocalendar()[:2] - ): - should_reset = True - elif key.balance_limit_reset == "monthly": - dt_now = datetime.fromtimestamp(now) - dt_reset = datetime.fromtimestamp(reset_date) - if dt_now.year > dt_reset.year or dt_now.month > dt_reset.month: - should_reset = True - - if should_reset: - key.total_spent = 0 - key.balance_limit_reset_date = now - session.add(key) - updated_count += 1 - - if updated_count > 0: - await session.commit() - logger.info( - "Periodic key reset complete", - extra={"keys_reset": updated_count}, - ) - - except asyncio.CancelledError: - break - except Exception as e: - logger.error(f"Error in periodic_key_reset: {e}") - - async def periodic_dead_key_prune() -> None: """Periodically prune dead API keys. Interval <= 0 disables it. diff --git a/routstr/balance.py b/routstr/balance.py index e7d4468d..f887f5ae 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -1,16 +1,13 @@ -import asyncio import hashlib -import time -from time import monotonic from typing import Annotated, NoReturn from fastapi import APIRouter, Depends, Header, HTTPException from fastapi.responses import JSONResponse from pydantic import BaseModel -from sqlmodel import col, select, update +from sqlmodel import col, select +from . import refund from .auth import ( - get_billing_key, redemption_error_to_http_exception, validate_bearer_key, ) @@ -21,20 +18,13 @@ from .core.db import ( get_session, release_stale_reservations, ) -from .core.db import ( - store_cashu_transaction_with_retry as store_cashu_transaction, -) from .core.logging import get_logger from .core.settings import settings from .lightning import lightning_router -from .payment.lnurl import MeltOutcomeAmbiguousError from .wallet import ( classify_redemption_error, credit_balance, - is_mint_connection_error, recieve_token, - send_to_lnurl, - send_token, token_mint_url, ) @@ -58,39 +48,15 @@ async def get_key_from_header( async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict: - billing_key = await get_billing_key(key, session) info = { "api_key": "sk-" + key.hashed_key, - "balance": billing_key.total_balance, - "reserved": billing_key.reserved_balance, - "is_child": key.parent_key_hash is not None, + "balance": key.total_balance, + "reserved": key.reserved_balance, "total_requests": key.total_requests, "total_spent": key.total_spent, - "balance_limit": key.balance_limit, - "balance_limit_reset": key.balance_limit_reset, "validity_date": key.validity_date, } - if key.parent_key_hash: - info["parent_key_preview"] = key.parent_key_hash[:8] + "..." - else: - # Fetch child keys if this is a parent key - statement = select(ApiKey).where(ApiKey.parent_key_hash == key.hashed_key) - results = await session.exec(statement) - child_keys = results.all() - if child_keys: - info["child_keys"] = [ - { - "api_key": "sk-" + ck.hashed_key, - "total_requests": ck.total_requests, - "total_spent": ck.total_spent, - "balance_limit": ck.balance_limit, - "balance_limit_reset": ck.balance_limit_reset, - "validity_date": ck.validity_date, - } - for ck in child_keys - ] - return info @@ -117,26 +83,18 @@ async def account_info( class BalanceCreateRequest(BaseModel): initial_balance_token: str - balance_limit: int | None = None - balance_limit_reset: str | None = None validity_date: int | None = None async def _create_balance( initial_balance_token: str, - balance_limit: int | None, - balance_limit_reset: str | None, validity_date: int | None, session: AsyncSession, ) -> dict: key = await validate_bearer_key(initial_balance_token, session) - if balance_limit is not None or balance_limit_reset or validity_date: - key.balance_limit = balance_limit - key.balance_limit_reset = balance_limit_reset + if validity_date is not None: key.validity_date = validity_date - if balance_limit_reset: - key.balance_limit_reset_date = int(time.time()) session.add(key) await session.commit() await session.refresh(key) @@ -154,8 +112,6 @@ async def create_balance_from_body( ) -> dict: return await _create_balance( payload.initial_balance_token, - payload.balance_limit, - payload.balance_limit_reset, payload.validity_date, session, ) @@ -164,15 +120,11 @@ async def create_balance_from_body( @router.get("/create") async def create_balance( initial_balance_token: str, - balance_limit: int | None = None, - balance_limit_reset: str | None = None, validity_date: int | None = None, session: AsyncSession = Depends(get_session), ) -> dict: return await _create_balance( initial_balance_token, - balance_limit, - balance_limit_reset, validity_date, session, ) @@ -208,7 +160,7 @@ async def topup_wallet_endpoint( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict[str, int]: - billing_key = await get_billing_key(key, session) + billing_key = key if topup_request is not None: cashu_token = topup_request.cashu_token @@ -276,35 +228,6 @@ async def topup_wallet_endpoint( return {"msats": amount_msats} -_REFUND_CACHE_TTL_SECONDS: int = settings.refund_cache_ttl_seconds -_refund_cache_lock: asyncio.Lock = asyncio.Lock() -_refund_cache: dict[str, tuple[float, dict[str, str]]] = {} - - -def _cache_key_for_authorization(authorization: str) -> str: - return hashlib.sha256(authorization.strip().encode()).hexdigest() - - -async def _refund_cache_get(authorization: str) -> dict[str, str] | None: - key = _cache_key_for_authorization(authorization) - async with _refund_cache_lock: - item = _refund_cache.get(key) - if item is None: - return None - expires_at, value = item - if expires_at <= monotonic(): - del _refund_cache[key] - return None - return value - - -async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None: - key = _cache_key_for_authorization(authorization) - expiry = monotonic() + _REFUND_CACHE_TTL_SECONDS - async with _refund_cache_lock: - _refund_cache[key] = (expiry, value) - - async def _lookup_key_no_create( bearer_value: str, session: AsyncSession ) -> ApiKey | None: @@ -318,17 +241,16 @@ async def _lookup_key_no_create( async def _get_persisted_api_key_refund( - key: ApiKey, session: AsyncSession + key: ApiKey, session: AsyncSession, token: str | None = None ) -> dict[str, str] | None: - result = await session.exec( - select(CashuTransaction) - .where( - CashuTransaction.api_key_hashed_key == key.hashed_key, - CashuTransaction.type == "out", - CashuTransaction.source == "apikey", - ) - .order_by(col(CashuTransaction.created_at).desc()) + query = select(CashuTransaction).where( + CashuTransaction.api_key_hashed_key == key.hashed_key, + CashuTransaction.type == "out", + CashuTransaction.source == "apikey", ) + if token is not None: + query = query.where(CashuTransaction.token == token) + result = await session.exec(query.order_by(col(CashuTransaction.created_at).desc())) refund = result.first() if refund is None: return None @@ -347,36 +269,13 @@ async def _get_persisted_api_key_refund( return persisted -async def _restore_balance( - session: AsyncSession, - hashed_key: str, - balance: int, - reserved_balance: int, - mint_url: str, -) -> None: - """Restore balance after a failed refund mint attempt.""" - restore_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == hashed_key) - .values( - balance=col(ApiKey.balance) + balance, - reserved_balance=col(ApiKey.reserved_balance) + reserved_balance, - ) - ) - await session.exec(restore_stmt) # type: ignore[call-overload] - await session.commit() - logger.info( - "refund_wallet_endpoint: balance restored after mint failure", - extra={ - "hashed_key": hashed_key, - "restored_balance": balance, - "mint_url": mint_url, - }, - ) +class RefundRequest(BaseModel): + lightning_address: str | None = None @router.post("/refund", response_model=None) async def refund_wallet_endpoint( + refund_request: RefundRequest | None = None, authorization: Annotated[str | None, Header()] = None, x_cashu: Annotated[str | None, Header()] = None, session: AsyncSession = Depends(get_session), @@ -452,17 +351,28 @@ async def refund_wallet_endpoint( }, ) + # Check for an open claim before any replay or destination lookup. + if open_claim := await refund.latest_open(session, key): + raise refund.refund_in_progress_error(open_claim) + if key.total_balance <= 0: - if cached := await _refund_cache_get(bearer_value): - return cached + paid = await refund.latest_terminal(session, key) + if paid and paid.method == "lightning": + return refund.describe(paid) + if paid and paid.token: + # Match the ledger row to this claim's token, not the latest one. + if persisted := await _get_persisted_api_key_refund( + key, session, paid.token + ): + return persisted + return refund.describe(paid) + # Legacy payouts predate the claim row, so fall back to the ledger. if persisted := await _get_persisted_api_key_refund(key, session): return persisted - - if key.parent_key_hash: - raise HTTPException( - status_code=400, - detail="Cannot refund child key. Please refund the parent key instead.", - ) + if paid: + return refund.describe(paid) + if stuck := await refund.latest_stuck(session, key): + raise refund.refund_in_progress_error(stuck) if key.reserved_balance > 0: # Release only durable reservations old enough to be stale. A newer @@ -481,168 +391,33 @@ async def refund_wallet_endpoint( logger.warning( "refund_wallet_endpoint: released stale reservation before refund", extra={ - "hashed_key": key.hashed_key, + "key_hash": key.hashed_key[:8], "stale_timeout_seconds": settings.stale_reservation_timeout_seconds, }, ) remaining_balance_msats: int = key.total_balance - - if key.refund_currency == "sat": - remaining_balance = remaining_balance_msats // 1000 - else: - remaining_balance = remaining_balance_msats + unit = refund.refund_unit(key) + remaining_balance = refund.amount_in_unit(remaining_balance_msats, unit) if remaining_balance_msats > 0 and remaining_balance <= 0: raise HTTPException(status_code=400, detail="Balance too small to refund") elif remaining_balance <= 0: raise HTTPException(status_code=400, detail="No balance to refund") - # Capture values before debit — the session may refresh key after commit - pre_debit_balance = key.balance - pre_debit_reserved = key.reserved_balance + requested = refund_request.lightning_address if refund_request else None + destination = requested or key.refund_address + if destination: + # Stored addresses can rot too; reject before any balance is debited. + await refund.validate_lightning_destination(destination) - # --- DEBIT FIRST: atomically zero the balance before minting tokens --- - # This prevents the race where a concurrent topup/spend happens between - # reading the balance and minting the refund token (double-spend). - debit_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.balance) == pre_debit_balance) - .where(col(ApiKey.reserved_balance) == pre_debit_reserved) - .values(balance=0, reserved_balance=0, reserved_at=None) + claim = await refund.open_claim( + session, + key, + method="lightning" if destination else "cashu", + destination=destination, ) - debit_result = await session.exec(debit_stmt) # type: ignore[call-overload] - await session.commit() - - if debit_result.rowcount == 0: - # Balance changed between read and debit — another request is active - raise HTTPException( - status_code=409, - detail="Balance changed concurrently. Please retry the refund.", - ) - - # The balance is locked at zero, so it is safe to create the refund token. - effective_refund_mint = ( - key.refund_mint_url - if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints - else settings.primary_mint - ) - try: - refund_currency = key.refund_currency or "sat" - if key.refund_address: - await send_to_lnurl( - remaining_balance, - key.refund_currency or "sat", - effective_refund_mint, - key.refund_address, - ) - result = {"recipient": key.refund_address} - else: - token = await send_token( - remaining_balance, refund_currency, effective_refund_mint - ) - effective_refund_mint = token_mint_url(token, effective_refund_mint) - result = {"token": token} - - if key.refund_currency == "sat": - result["sats"] = str(remaining_balance_msats // 1000) - else: - result["msats"] = str(remaining_balance_msats) - - if "token" in result: - logger.info( - "refund_wallet_endpoint: cashu token issued", - extra={ - "path": "/v1/wallet/refund", - "token": result["token"], - "amount": remaining_balance, - "currency": key.refund_currency or "sat", - }, - ) - - except MeltOutcomeAmbiguousError as e: - # The melt was dispatched and may still settle. Restoring the balance - # here would let the same debit be paid out twice; keep the debit and - # leave the outcome to reconciliation. - logger.error( - "refund_wallet_endpoint: melt outcome ambiguous; balance withheld " - "pending reconciliation", - extra={ - "error": str(e), - "hashed_key": key.hashed_key, - "remaining_balance": remaining_balance, - "refund_currency": key.refund_currency, - "refund_mint_url": key.refund_mint_url, - }, - ) - raise HTTPException( - status_code=502, - detail=( - "Refund was dispatched but its outcome is unconfirmed; the " - "balance is withheld until reconciliation completes" - ), - ) - except HTTPException: - # Minting failed — restore the debited balance - await _restore_balance( - session, - key.hashed_key, - pre_debit_balance, - pre_debit_reserved, - key.refund_mint_url or "", - ) - raise - except Exception as e: - # Minting failed — restore the debited balance - await _restore_balance( - session, - key.hashed_key, - pre_debit_balance, - pre_debit_reserved, - key.refund_mint_url or "", - ) - error_msg = str(e) - logger.error( - "refund_wallet_endpoint: mint/send failed", - extra={ - "error": error_msg, - "error_type": type(e).__name__, - "hashed_key": key.hashed_key, - "remaining_balance": remaining_balance, - "refund_currency": key.refund_currency, - "refund_mint_url": key.refund_mint_url, - "has_refund_address": bool(key.refund_address), - }, - ) - if is_mint_connection_error(e): - raise HTTPException(status_code=503, detail="Mint service unavailable") - else: - raise HTTPException(status_code=500, detail="Refund failed") - - await _refund_cache_set(bearer_value, result) - - if "token" in result: - await store_cashu_transaction( - token=result["token"], - amount=remaining_balance, - unit=key.refund_currency or "sat", - mint_url=effective_refund_mint, - typ="out", - collected=False, - source="apikey", - api_key_hashed_key=key.hashed_key, - ) - - logger.info( - "refund_wallet_endpoint: refund successful", - extra={ - "refunded_msats": remaining_balance_msats, - "previous_reserved_balance": key.reserved_balance, - }, - ) - - return result + return await refund.execute(session, claim) @router.get("/history") @@ -650,12 +425,6 @@ async def wallet_history( key: ApiKey = Depends(get_key_from_header), session: AsyncSession = Depends(get_session), ) -> dict[str, list[dict[str, str | int | bool | None]]]: - if key.parent_key_hash: - raise HTTPException( - status_code=400, - detail="Cannot view child key history. Please use the parent key instead.", - ) - result = await session.exec( select(CashuTransaction) .where(CashuTransaction.api_key_hashed_key == key.hashed_key) @@ -693,132 +462,6 @@ async def donate(token: str, ref: str | None = None) -> str: return "Invalid token." -class ChildKeyRequest(BaseModel): - count: int - balance_limit: int | None = None - balance_limit_reset: str | None = None - validity_date: int | None = None - - -@router.post("/child-key") -async def create_child_key( - payload: ChildKeyRequest, - key: ApiKey = Depends(get_key_from_header), - session: AsyncSession = Depends(get_session), -) -> dict: - """Creates one or more child API keys that use the parent's balance.""" - # Log incoming request for debugging - logger.debug(f"Child key creation request: count={payload.count}") - - count = payload.count - if count < 1 or count > 50: - raise HTTPException(status_code=400, detail="Count must be between 1 and 50.") - - # Check if this is already a child key - if key.parent_key_hash: - raise HTTPException( - status_code=400, - detail="Cannot create a child key for another child key.", - ) - - cost_per_key = settings.child_key_cost - total_cost = cost_per_key * count - - if key.total_balance < total_cost: - raise HTTPException( - status_code=402, - detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.", - ) - - # Deduct cost from parent atomically — guards against concurrent requests - # that both pass the balance check above on stale in-memory state. - deduct_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= total_cost) - .values( - balance=col(ApiKey.balance) - total_cost, - total_spent=col(ApiKey.total_spent) + total_cost, - ) - ) - result = await session.exec(deduct_stmt) # type: ignore[call-overload] - - if result.rowcount == 0: - raise HTTPException( - status_code=402, - detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.", - ) - - # Generate new keys - import secrets - - new_keys = [] - for _ in range(count): - new_key_raw = secrets.token_hex(32) - new_key_hash = new_key_raw # We use the raw key as the hash for sk- keys - - child_key = ApiKey( - hashed_key=new_key_hash, - balance=0, - parent_key_hash=key.hashed_key, - balance_limit=payload.balance_limit, - balance_limit_reset=payload.balance_limit_reset, - balance_limit_reset_date=int(time.time()) - if payload.balance_limit_reset - else None, - validity_date=payload.validity_date, - ) - session.add(child_key) - new_keys.append("sk-" + new_key_hash) - - await session.commit() - await session.refresh(key) - - response_data = { - "api_keys": new_keys, - "count": count, - "cost_msats": total_cost, - "cost_sats": total_cost // 1000, - "parent_balance": key.balance, - "parent_balance_sats": key.balance // 1000, - } - logger.debug(f"Child key creation response: {response_data}") - return response_data - - -class ChildKeyResetRequest(BaseModel): - child_key: str - - -@router.post("/child-key/reset") -async def reset_child_key_spent( - payload: ChildKeyResetRequest, - key: ApiKey = Depends(get_key_from_header), - session: AsyncSession = Depends(get_session), -) -> dict: - """Resets the total_spent of a child key. Must be called by the parent.""" - child_key_raw = payload.child_key - if child_key_raw.startswith("sk-"): - child_key_raw = child_key_raw[3:] - - child_key = await session.get(ApiKey, child_key_raw) - if not child_key: - raise HTTPException(status_code=404, detail="Child key not found.") - - if child_key.parent_key_hash != key.hashed_key: - raise HTTPException( - status_code=403, detail="Unauthorized. You are not the parent of this key." - ) - - child_key.total_spent = 0 - if child_key.balance_limit_reset: - child_key.balance_limit_reset_date = int(time.time()) - session.add(child_key) - await session.commit() - - return {"success": True, "message": "Child key balance reset successfully."} - - @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/cashu_compat.py b/routstr/cashu_compat.py new file mode 100644 index 00000000..02a56ca4 --- /dev/null +++ b/routstr/cashu_compat.py @@ -0,0 +1,108 @@ +"""Compatibility shim that keeps cashu 0.20.x working on httpx>=0.28. + +cashu's ``async_set_httpx_client`` decorator builds the client for *every* mint +call as ``httpx.AsyncClient(proxies=proxies_dict, ...)``. httpx deprecated +``proxies`` in 0.26 and removed it in 0.28, so on httpx>=0.28 every wallet +operation routstr performs -- ``load_mint_keysets``, ``mint_quote``, +``melt_quote``, token redeem -- raises:: + + TypeError: AsyncClient.__init__() got an unexpected keyword argument 'proxies' + +Upstream cashu (0.20.3, the latest release) still caps ``httpx<0.26`` and has no +release that fixes this, while litellm>=1.84 requires ``httpx>=0.28``. We bridge +the gap by translating the keyword *inside cashu's module namespace only* -- +global httpx behaviour is untouched, and the shim is a no-op on httpx<0.28. + +Delete this module once cashu ships a release that passes ``proxy=``. +""" + +from typing import Any + +import httpx + +__all__ = ["install_cashu_httpx_shim"] + +_ALL_SCHEMES = "all://" + + +def _single_proxy(proxies: Any) -> str | None: + """Collapse an httpx<0.28 ``proxies`` mapping into a single ``proxy`` URL. + + cashu only ever builds ``{}`` or ``{"all://": url}``, so a mapping with one + distinct URL is all we need to support. Anything richer is unrepresentable + as httpx 0.28's scalar ``proxy=``; we fail closed and raise rather than + return ``None``, because dropping the entry would silently send mint + traffic direct instead of through the configured Tor/SOCKS proxy. + """ + if not proxies: + return None + if isinstance(proxies, str): + return proxies + if isinstance(proxies, dict): + if _ALL_SCHEMES in proxies: + value = proxies[_ALL_SCHEMES] + return str(value) if value is not None else None + distinct = {str(v) for v in proxies.values() if v is not None} + if not distinct: + return None + if len(distinct) == 1: + return distinct.pop() + raise ValueError( + f"cannot represent proxies={proxies!r} as httpx 0.28 proxy=; " + "refusing to send proxied traffic direct" + ) + + +class _ProxiesCompatAsyncClient(httpx.AsyncClient): + """``httpx.AsyncClient`` that still accepts the removed ``proxies`` kwarg.""" + + def __init__(self, *args: Any, **kwargs: Any) -> None: + if "proxies" in kwargs: + proxies = kwargs.pop("proxies") + proxy = _single_proxy(proxies) + if proxy is not None and kwargs.get("proxy") is None: + kwargs["proxy"] = proxy + elif isinstance(proxies, dict) and not proxies: + # In httpx<0.28, proxies={} disabled environment proxy discovery. + kwargs.setdefault("trust_env", False) + super().__init__(*args, **kwargs) + + +class _HttpxNamespace: + """Stand-in for the ``httpx`` module inside cashu's ``v1_api``. + + Every attribute resolves against the real module except ``AsyncClient``, + so cashu keeps using genuine httpx types everywhere else. + """ + + AsyncClient = _ProxiesCompatAsyncClient + + def __getattr__(self, name: str) -> Any: + return getattr(httpx, name) + + +def _httpx_accepts_proxies() -> bool: + import inspect + + try: + return "proxies" in inspect.signature(httpx.AsyncClient.__init__).parameters + except (TypeError, ValueError): # pragma: no cover - defensive + return False + + +def install_cashu_httpx_shim() -> bool: + """Patch cashu's mint client to survive httpx>=0.28. + + Returns True when the shim was installed, False when it wasn't needed. + Safe to call repeatedly. + """ + if _httpx_accepts_proxies(): + return False + + from cashu.wallet import v1_api + + if isinstance(getattr(v1_api, "httpx", None), _HttpxNamespace): + return True + + v1_api.httpx = _HttpxNamespace() # type: ignore[assignment] + return True diff --git a/routstr/core/admin.py b/routstr/core/admin.py index c314ae14..cc140f6c 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1,4 +1,3 @@ -import asyncio import json import re import secrets @@ -6,12 +5,17 @@ from datetime import datetime, timezone from pathlib import Path from fastapi import APIRouter, Depends, HTTPException, Query, Request -from pydantic import BaseModel, RootModel +from pydantic import BaseModel, RootModel, field_validator from pydantic.v1 import ValidationError as PydanticValidationError from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession -from ..payment.models import _row_to_model, list_models +from ..payment.models import ( + REQUIRED_PRICING_FIELDS, + _row_to_model, + list_models, +) +from ..payment.rates import BILLABLE_PRICING_FIELDS, coerce_rate from ..proxy import refresh_model_maps, reinitialize_upstreams from ..wallet import fetch_all_balances, send_token, token_mint_url from . import vault @@ -30,6 +34,7 @@ from .db import ( from .db import ( store_cashu_transaction_with_retry as store_cashu_transaction, ) +from .exceptions import json_compliant from .log_manager import log_manager from .logging import get_logger from .provider_slugs import allocate_unique_provider_slug @@ -69,7 +74,9 @@ async def require_admin_api(request: Request) -> None: async with create_session() as session: result = await session.exec(select(CliToken).where(CliToken.token == token)) cli_token = result.first() - if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts): + if cli_token and ( + cli_token.expires_at is None or cli_token.expires_at > now_ts + ): cli_token.last_used_at = now_ts session.add(cli_token) await session.commit() @@ -104,25 +111,28 @@ async def get_temporary_balances_api( ) total = count_result.one() - # Aggregate totals across the whole (search-filtered) set, not just the - # current page. Balance counts only parent (non-child) keys to avoid - # double-counting, since child keys draw from their parent's balance. - totals_result = await session.exec( + # Aggregate totals across the whole search-filtered set, not just this page. + balance_totals_result = await session.exec( select( + func.coalesce(func.sum(ApiKey.balance), 0), + func.coalesce(func.sum(ApiKey.reserved_balance), 0), func.coalesce( - func.sum( - case( - (col(ApiKey.parent_key_hash).is_(None), ApiKey.balance), - else_=0, - ) - ), - 0, + func.sum(col(ApiKey.balance) - col(ApiKey.reserved_balance)), 0 ), + ).where(*filters) + ) + ( + total_balance, + total_reserved_balance, + total_available_balance, + ) = balance_totals_result.one() + usage_totals_result = await session.exec( + select( func.coalesce(func.sum(ApiKey.total_spent), 0), func.coalesce(func.sum(ApiKey.total_requests), 0), ).where(*filters) ) - total_balance, total_spent, total_requests = totals_result.one() + total_spent, total_requests = usage_totals_result.one() # Latest created first; keys with no created_at (legacy rows) sort last. # Use an explicit CASE rather than relying on dialect NULL-ordering so @@ -143,13 +153,12 @@ async def get_temporary_balances_api( { "hashed_key": key.hashed_key, "balance": key.balance, + "reserved_balance": key.reserved_balance, + "available_balance": key.total_balance, "total_spent": key.total_spent, "total_requests": key.total_requests, "refund_address": key.refund_address, "key_expiry_time": key.key_expiry_time, - "parent_key_hash": key.parent_key_hash, - "balance_limit": key.balance_limit, - "balance_limit_reset": key.balance_limit_reset, "validity_date": key.validity_date, "created_at": key.created_at, } @@ -158,6 +167,8 @@ async def get_temporary_balances_api( "total": total, "totals": { "total_balance": total_balance, + "total_reserved_balance": total_reserved_balance, + "total_available_balance": total_available_balance, "total_spent": total_spent, "total_requests": total_requests, }, @@ -165,8 +176,6 @@ async def get_temporary_balances_api( class ApiKeyUpdate(BaseModel): - balance_limit: int | None = None - balance_limit_reset: str | None = None validity_date: int | None = None @@ -181,10 +190,6 @@ async def update_apikey( if not key: raise HTTPException(status_code=404, detail="API key not found") - if update.balance_limit is not None: - key.balance_limit = update.balance_limit - if update.balance_limit_reset is not None: - key.balance_limit_reset = update.balance_limit_reset if update.validity_date is not None: key.validity_date = update.validity_date @@ -194,8 +199,6 @@ async def update_apikey( return { "hashed_key": key.hashed_key, - "balance_limit": key.balance_limit, - "balance_limit_reset": key.balance_limit_reset, "validity_date": key.validity_date, } @@ -256,16 +259,12 @@ async def update_password(request: Request, password_update: PasswordUpdate) -> secret = await get_secret(session) if not secret.admin_password_hash: - raise HTTPException( - status_code=500, detail="Admin password not configured" - ) + raise HTTPException(status_code=500, detail="Admin password not configured") if not vault.verify_password( password_update.current_password, secret.admin_password_hash ): - raise HTTPException( - status_code=401, detail="Current password is incorrect" - ) + raise HTTPException(status_code=401, detail="Current password is incorrect") # Validate new password new_password = password_update.new_password.strip() @@ -483,6 +482,38 @@ class ModelCreate(BaseModel): enabled: bool = True forwarded_model_id: str | None = None + @field_validator("pricing") + @classmethod + def _validate_pricing(cls, value: dict[str, object]) -> dict[str, object]: + """Reject a rate that is malformed, non-finite, negative or not there. + + A present-but-invalid rate would otherwise slip through: a non-numeric + string coerces to $0 on the read path (an unpriced-looking row), while a + negative or ``NaN``/``inf`` value is truthy and reads back as a real + price, so the model could be enabled and bill a nonsensical amount. + Surfacing a 422 reports the client bug as a client bug instead of + persisting it. Numeric strings (``"0.000005"``) stay valid, and so does + an omitted auxiliary rate — the stored JSON accepts both. + """ + for field in BILLABLE_PRICING_FIELDS: + if field not in value: + # ``dict.get`` cannot tell this from an explicit ``null``, so + # both were skipped and a row that ``Pricing`` cannot parse was + # written — and then raised out of the response that reads it + # back, after the row had been committed. + if field in REQUIRED_PRICING_FIELDS: + raise ValueError(f"{field} is required") + continue + # The shared coercion also absorbs the OverflowError an oversized + # integer raises, which pydantic does not convert into a validation + # error — unhandled it escaped as a 500 for a bad client value. + if coerce_rate(value[field]) is None: + raise ValueError( + f"{field} must be a finite, non-negative number, " + f"got {value[field]!r}" + ) + return value + def _normalize_forwarded_model_id(value: str | None) -> str | None: if value is None: @@ -609,9 +640,13 @@ async def get_provider_model(provider_id: str, model_id: str) -> dict[str, objec raise HTTPException( status_code=404, detail="Model not found for this provider" ) - return _row_to_model( - row, apply_provider_fee=False, provider_fee=provider.provider_fee - ).dict() # type: ignore + # Same duty as the listing this view is opened from: a stored rate that + # is not a usable number must be shown as it is, not encoded as `null`. + return json_compliant( # type: ignore[return-value] + _row_to_model( + row, apply_provider_fee=False, provider_fee=provider.provider_fee + ).dict() + ) @admin_router.delete( @@ -870,29 +905,43 @@ class UpstreamProviderUpdateBySlug(BaseModel): provider_settings: dict | None = None -async def _active_ppq_claim_in_session(session: AsyncSession, provider_id: int) -> bool: +async def _active_auto_topup_claim_in_session( + session: AsyncSession, provider_id: int, provider_type: str +) -> bool: """Check for an active claim inside the caller's transaction. Must share the transaction of whatever destructive write it is guarding — a check in its own session leaves a window for a worker to create the claim between the check and the commit. """ - from ..upstream.auto_topup import _ppq_state_id_for_provider + from ..upstream.auto_topup import ( + _ppq_state_id_for_provider, + _routstr_state_id_for_provider, + ) - claim = await session.get(CashuTransaction, _ppq_state_id_for_provider(provider_id)) + state_id = ( + _ppq_state_id_for_provider(provider_id) + if provider_type == "ppqai" + else _routstr_state_id_for_provider(provider_id) + ) + claim = await session.get(CashuTransaction, state_id) return claim is not None and not claim.collected and not claim.swept -def _require_valid_ppq_auto_topup( - provider_type: str, settings: dict | None -) -> None: - """Reject PPQ auto top-up settings the worker would later refuse.""" - if provider_type != "ppqai": +def _require_valid_auto_topup(provider_type: str, settings: dict | None) -> None: + """Reject auto top-up settings the worker would later refuse.""" + from ..upstream.auto_topup import ( + validate_ppq_auto_topup_settings, + validate_routstr_auto_topup_settings, + ) + + if provider_type == "ppqai": + problem = validate_ppq_auto_topup_settings(settings) + elif provider_type == "routstr": + problem = validate_routstr_auto_topup_settings(settings) + else: return - from ..upstream.auto_topup import validate_ppq_auto_topup_settings - - problem = validate_ppq_auto_topup_settings(settings) if problem is not None: raise HTTPException(status_code=400, detail=problem) @@ -917,16 +966,18 @@ async def _apply_provider_update( ) if ( provider_type_changed - and provider.provider_type == "ppqai" + and provider.provider_type in ("ppqai", "routstr") and provider.id is not None - and await _active_ppq_claim_in_session(session, provider.id) + and await _active_auto_topup_claim_in_session( + session, provider.id, provider.provider_type + ) ): - # Changing the type would orphan the claim: the PPQ endpoints refuse - # non-ppqai providers, so nobody could ever inspect or release it. + # Changing the type would orphan the claim: the claim endpoints refuse + # providers of the wrong type, so nobody could inspect or release it. raise HTTPException( status_code=409, detail=( - "This provider has an active PPQ auto top-up claim. Release " + "This provider has an active auto top-up claim. Release " "it before changing the provider type" ), ) @@ -977,7 +1028,7 @@ async def _apply_provider_update( except (json.JSONDecodeError, TypeError): effective_settings = None if effective_settings is not None: - _require_valid_ppq_auto_topup(provider.provider_type, effective_settings) + _require_valid_auto_topup(provider.provider_type, effective_settings) if payload.provider_settings is not None: provider.provider_settings = json.dumps(payload.provider_settings) @@ -1017,9 +1068,7 @@ async def create_upstream_provider( else: slug = await allocate_unique_provider_slug(session, payload.provider_type) - _require_valid_ppq_auto_topup( - payload.provider_type, payload.provider_settings - ) + _require_valid_auto_topup(payload.provider_type, payload.provider_settings) provider = UpstreamProviderRow( slug=slug, @@ -1078,9 +1127,7 @@ async def update_upstream_provider_by_slug( lookup = _validate_slug(payload.slug) async with create_session() as session: result = await session.exec( - select(UpstreamProviderRow).where( - UpstreamProviderRow.slug == lookup - ) + select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup) ) provider = result.first() if not provider: @@ -1117,16 +1164,19 @@ async def delete_upstream_provider(provider_id: str) -> dict[str, object]: # re-reads the provider inside its own transaction, so these two # writes serialise — either the claim lands first and this 409s, or # the delete lands first and the worker refuses to claim. - if provider.provider_type == "ppqai" and await _active_ppq_claim_in_session( - session, deleted_id + if provider.provider_type in ( + "ppqai", + "routstr", + ) and await _active_auto_topup_claim_in_session( + session, deleted_id, provider.provider_type ): # Deleting now would orphan the claim and any funds it tracks: - # the PPQ endpoints 404 without the provider row, so the claim + # the claim endpoints 404 without the provider row, so the claim # could never again be inspected or released. raise HTTPException( status_code=409, detail=( - "This provider has an active PPQ auto top-up claim. " + "This provider has an active auto top-up claim. " "Resolve and release it before deleting the provider" ), ) @@ -1186,8 +1236,13 @@ async def get_provider_models(provider_id: str) -> dict[str, object]: "provider_type": provider.provider_type, "base_url": provider.base_url, }, - "db_models": [m.dict() for m in db_models], - "remote_models": [m.dict() for m in filtered_remote_models], + # This listing includes disabled models, so it is the one view that + # still carries a row the served-catalog backstop holds back — + # including one whose stored rate is not a usable number. The + # encoder would report that rate as `null`, indistinguishable from a + # missing one; show the operator the value that needs fixing. + "db_models": [json_compliant(m.dict()) for m in db_models], + "remote_models": [json_compliant(m.dict()) for m in filtered_remote_models], } @@ -1316,50 +1371,37 @@ async def initiate_provider_topup( else {} ) - last_status_code = 500 - last_error_detail: object = "Failed to create top-up invoice" + # Quote creation is unsafe to retry without idempotency. + resp = await client.post( + f"{clean_url}/v1/balance/lightning/invoice", + json=request_json, + headers=headers, + ) - # Some upstream Routstr nodes fail the first invoice request after warm-up - # and succeed immediately on retry. Retry once here so the UI stays single-click. - for attempt in range(2): - resp = await client.post( - f"{clean_url}/v1/balance/lightning/invoice", - json=request_json, - headers=headers, - ) - - if resp.status_code == 200: - data = resp.json() - return { - "ok": True, - "topup_data": { - "payment_request": data.get("bolt11"), - "invoice_id": data.get("invoice_id"), - "status": "pending", - }, - } - - logger.error( - f"Upstream topup request failed: {resp.text}", - extra={ - "provider_id": provider_id, - "attempt": attempt + 1, - "status_code": resp.status_code, + if resp.status_code == 200: + data = resp.json() + return { + "ok": True, + "topup_data": { + "payment_request": data.get("bolt11"), + "invoice_id": data.get("invoice_id"), + "status": "pending", }, - ) - try: - last_error_detail = resp.json() - except Exception: - last_error_detail = resp.text - last_status_code = resp.status_code - - if resp.status_code < 500 or attempt == 1: - break - - await asyncio.sleep(0.2) + } + logger.error( + f"Upstream topup request failed: {resp.text}", + extra={ + "provider_id": provider_id, + "status_code": resp.status_code, + }, + ) + try: + error_detail: object = resp.json() + except Exception: + error_detail = resp.text raise HTTPException( - status_code=last_status_code, detail=last_error_detail + status_code=resp.status_code, detail=error_detail ) upstream_instance = _instantiate_provider(provider) @@ -1718,6 +1760,34 @@ async def get_logs_api( } +@admin_router.get( + "/api/logs/request/{request_id}", dependencies=[Depends(require_admin_api)] +) +async def get_logs_by_request_id_api( + request: Request, + request_id: str, + date: str | None = None, + limit: int = Query(default=200, ge=1, le=1000), +) -> dict[str, object]: + """ + Get every log entry belonging to a single request ID, oldest first. + """ + log_entries = log_manager.search_logs( + date=date, + request_id=request_id, + limit=limit, + ) + log_entries.sort(key=lambda entry: str(entry.get("asctime", ""))) + + return { + "logs": log_entries, + "total": len(log_entries), + "request_id": request_id, + "date": date, + "limit": limit, + } + + @admin_router.get("/api/logs/dates", dependencies=[Depends(require_admin_api)]) async def get_log_dates_api(request: Request) -> dict[str, object]: logs_dir = Path("logs") @@ -1811,6 +1881,87 @@ async def release_ppq_auto_topup_api( return {"ok": True, "released": True} +_ROUTSTR_RELEASE_ERRORS = { + "no_active_claim": "No active Routstr auto top-up claim to release", + "stale_state": "The claim changed since it was reviewed; reload and check again", + "claim_changed": ( + "The claim changed while the release was being applied; reload and check again" + ), +} + + +class ReleaseRoutstrAutoTopupRequest(BaseModel): + confirmed_peer_reconciled: bool + state_token: str | None = None + + +async def _require_routstr_provider(provider_id: int) -> UpstreamProviderRow: + async with create_session() as session: + provider = await session.get(UpstreamProviderRow, provider_id) + if provider is None: + raise HTTPException(status_code=404, detail="Provider not found") + if provider.provider_type != "routstr": + raise HTTPException(status_code=400, detail="Provider is not a Routstr node") + return provider + + +@admin_router.get( + "/api/upstream-providers/{provider_id}/routstr-auto-topup", + dependencies=[Depends(require_admin_api)], +) +async def get_routstr_auto_topup_api(provider_id: int) -> dict[str, object]: + await _require_routstr_provider(provider_id) + from ..upstream.auto_topup import get_routstr_auto_topup_state + + return {"ok": True, **await get_routstr_auto_topup_state(provider_id)} + + +@admin_router.post( + "/api/upstream-providers/{provider_id}/routstr-auto-topup/release", + dependencies=[Depends(require_admin_api)], +) +async def release_routstr_auto_topup_api( + provider_id: int, payload: ReleaseRoutstrAutoTopupRequest +) -> dict[str, object]: + await _require_routstr_provider(provider_id) + if not payload.confirmed_peer_reconciled: + raise HTTPException( + status_code=400, + detail="Confirm the peer credited or returned the token before releasing", + ) + + from ..upstream.auto_topup import release_routstr_auto_topup_state + + outcome = await release_routstr_auto_topup_state( + provider_id, state_token=payload.state_token + ) + if not outcome.released: + raise HTTPException( + status_code=409, detail=_ROUTSTR_RELEASE_ERRORS[outcome.reason] + ) + + logger.warning( + "Admin released Routstr auto top-up claim after manual reconciliation", + extra={"provider_id": provider_id, "state_token": payload.state_token}, + ) + return {"ok": True, "released": True} + + +def _transaction_status(tx: CashuTransaction) -> str: + """An outgoing admin withdrawal ends at "issued": the node hands the bearer + token over and never learns whether it was redeemed, so its flags stay false + and it would otherwise read pending forever. Incoming admin rows are redeemed + by the node itself and keep the normal collected/swept lifecycle. + """ + if tx.swept: + return "swept" + if tx.collected: + return "collected" + if tx.source == "admin" and tx.type == "out": + return "issued" + return "pending" + + @admin_router.get("/api/transactions", dependencies=[Depends(require_admin_api)]) async def get_transactions_api( type: str | None = None, @@ -1823,10 +1974,11 @@ async def get_transactions_api( async with create_session() as session: from sqlmodel import col, func - # Hide only the deterministic PPQ claim-lock rows. Append-only PPQ - # payment rows remain visible as the audit trail for irreversible melts. + # Hide only the deterministic claim-lock rows. Append-only PPQ payment + # rows and auto-topup token rows remain visible as the audit trail. base = select(CashuTransaction).where( - ~col(CashuTransaction.id).like("ppq-auto-topup-%") + ~col(CashuTransaction.id).like("ppq-auto-topup-%"), + ~col(CashuTransaction.id).like("routstr-auto-topup-%"), ) if type: base = base.where(CashuTransaction.type == type) @@ -1840,13 +1992,25 @@ async def get_transactions_api( base = base.where(CashuTransaction.source == source) if status: if status == "collected": - base = base.where(CashuTransaction.collected == True) # noqa: E712 + base = base.where( + CashuTransaction.collected == True, # noqa: E712 + CashuTransaction.swept == False, # noqa: E712 + ) elif status == "swept": base = base.where(CashuTransaction.swept == True) # noqa: E712 + elif status == "issued": + base = base.where( + CashuTransaction.source == "admin", + CashuTransaction.type == "out", + CashuTransaction.collected == False, # noqa: E712 + CashuTransaction.swept == False, # noqa: E712 + ) elif status == "pending": base = base.where( CashuTransaction.collected == False, # noqa: E712 CashuTransaction.swept == False, # noqa: E712 + (CashuTransaction.source != "admin") + | (CashuTransaction.type != "out"), ) if search: @@ -1873,15 +2037,17 @@ async def get_transactions_api( return { "transactions": [ - tx.dict(exclude={"sweep_started_at"}) for tx in transactions + { + **tx.dict(exclude={"sweep_started_at"}), + "status": _transaction_status(tx), + } + for tx in transactions ], "total": total, } -@admin_router.get( - "/api/lightning-invoices", dependencies=[Depends(require_admin_api)] -) +@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)]) async def get_lightning_invoices_api( status: str | None = None, purpose: str | None = None, diff --git a/routstr/core/db.py b/routstr/core/db.py index f31691a1..eb4ba2fe 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -12,12 +12,11 @@ from typing import AsyncGenerator from alembic import command from alembic.config import Config from alembic.util.exc import CommandError -from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_ +from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_, text from sqlalchemy.engine import make_url from sqlalchemy.exc import IntegrityError, OperationalError from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio.engine import create_async_engine -from sqlalchemy.orm import aliased from sqlmodel import Field, Relationship, SQLModel, col, func, select, update from sqlmodel.ext.asyncio.session import AsyncSession @@ -37,6 +36,9 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine: is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"} pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping} + connect_args: dict[str, object] = {} + if is_sqlite and not is_memory_sqlite: + connect_args["timeout"] = settings.database_busy_timeout if not is_memory_sqlite: options.update( pool_size=settings.database_pool_size, @@ -51,9 +53,12 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine: "database_url_backend": backend, "in_memory_sqlite": is_memory_sqlite, **options, + "connect_args": connect_args, }, ) - created_engine = create_async_engine(database_url, echo=False, **options) + created_engine = create_async_engine( + database_url, echo=False, connect_args=connect_args, **options + ) hold_warn_seconds = settings.database_pool_hold_warn_seconds def record_pool_checkout( @@ -132,21 +137,6 @@ class ApiKey(SQLModel, table=True): # type: ignore default=None, description="Currency of the cashu-token", ) - parent_key_hash: str | None = Field( - default=None, foreign_key="api_keys.hashed_key", index=True - ) - balance_limit: int | None = Field( - default=None, - description="Max spendable balance in msats for this key (mostly for child keys)", - ) - balance_limit_reset: str | None = Field( - default=None, - description="Reset policy for balance limit (manual, daily, monthly, etc.)", - ) - balance_limit_reset_date: int | None = Field( - default=None, - description="Unix timestamp of the last time the balance limit was reset", - ) validity_date: int | None = Field( default=None, description="Unix timestamp after which the key is no longer valid", @@ -171,6 +161,55 @@ async def reset_all_reserved_balances(session: AsyncSession) -> None: logger.info("Reset reserved balances on startup") +async def _transition_stale_reservation( + session: AsyncSession, reservation_id: str, cutoff: int +) -> bool: + """Mark one reservation released iff its lease is still older than cutoff. + + ``created_at`` doubles as the heartbeat lease timestamp, so the guard must + be part of this update: a reservation renewed between the sweeper's select + and this transition is in flight and must survive. + """ + transition = await session.exec( # type: ignore[call-overload] + update(ReservationRelease) + .where(col(ReservationRelease.id) == reservation_id) + .where(col(ReservationRelease.status) == "active") + .where(col(ReservationRelease.created_at) < cutoff) + .values(status="released") + ) + return bool(transition.rowcount == 1) + + +async def _release_legacy_aggregate( + session: AsyncSession, + key_hash: str, + observed_reserved: int, + observed_reserved_at: int | None, +) -> bool: + """Zero one legacy aggregate reservation iff it is exactly as observed. + + A new reservation committing between the sweeper's read and this update + changes ``reserved_balance``/``reserved_at`` in the same transaction that + creates its durable row, so this compare-and-swap fails instead of erasing + the newcomer's reserved funds. + """ + if observed_reserved <= 0: + return False + reserved_at_guard = ( + col(ApiKey.reserved_at).is_(None) + if observed_reserved_at is None + else col(ApiKey.reserved_at) == observed_reserved_at + ) + result = await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == key_hash) + .where(col(ApiKey.reserved_balance) == observed_reserved) + .where(reserved_at_guard) + .values(reserved_balance=0, reserved_at=None) + ) + return bool(result.rowcount == 1) + + async def release_stale_reservations( session: AsyncSession, max_age_seconds: int, @@ -191,25 +230,22 @@ async def release_stale_reservations( col(ReservationRelease.billing_key_hash) == key_hash, ) ) - reservations = (await session.exec(query)).all() + # Capture primitives: a repair rollback below would expire ORM instances. + reservation_rows = [ + (r.id, r.key_hash, r.billing_key_hash, r.reserved_msats) + for r in (await session.exec(query)).all() + ] released = 0 - for reservation in reservations: - transition = await session.exec( # type: ignore[call-overload] - update(ReservationRelease) - .where(col(ReservationRelease.id) == reservation.id) - .where(col(ReservationRelease.status) == "active") - .values(status="released") - ) - if transition.rowcount != 1: + for res_id, res_key_hash, res_billing_hash, res_msats in reservation_rows: + if not await _transition_stale_reservation(session, res_id, cutoff): continue values = { - "reserved_balance": col(ApiKey.reserved_balance) - - reservation.reserved_msats, + "reserved_balance": col(ApiKey.reserved_balance) - res_msats, "reserved_at": case( ( - col(ApiKey.reserved_balance) - reservation.reserved_msats > 0, + col(ApiKey.reserved_balance) - res_msats > 0, col(ApiKey.reserved_at), ), else_=None, @@ -217,24 +253,42 @@ async def release_stale_reservations( } parent_result = await session.exec( # type: ignore[call-overload] update(ApiKey) - .where(col(ApiKey.hashed_key) == reservation.billing_key_hash) - .where(col(ApiKey.reserved_balance) >= reservation.reserved_msats) + .where(col(ApiKey.hashed_key) == res_billing_hash) + .where(col(ApiKey.reserved_balance) >= res_msats) .values(**values) ) - if parent_result.rowcount != 1: - await session.rollback() - return 0 - - if reservation.billing_key_hash != reservation.key_hash: + aggregates_ok = parent_result.rowcount == 1 + if aggregates_ok and res_billing_hash != res_key_hash: child_result = await session.exec( # type: ignore[call-overload] update(ApiKey) - .where(col(ApiKey.hashed_key) == reservation.key_hash) - .where(col(ApiKey.reserved_balance) >= reservation.reserved_msats) + .where(col(ApiKey.hashed_key) == res_key_hash) + .where(col(ApiKey.reserved_balance) >= res_msats) .values(**values) ) - if child_result.rowcount != 1: - await session.rollback() - return 0 + aggregates_ok = child_result.rowcount == 1 + + if not aggregates_ok: + # The aggregates no longer hold this reservation's msats — the + # durable row is corrupt. Repair by terminalizing it WITHOUT + # subtracting uncertain aggregates (legacy cleanup below reconciles + # any stale remainder) and keep sweeping the rest of the batch: + # one corrupt row must not poison all stale cleanup. + await session.rollback() + if await _transition_stale_reservation(session, res_id, cutoff): + await session.commit() + released += 1 + logger.error( + "Released corrupt stale reservation without aggregate subtraction", + extra={ + "reservation_id": res_id, + "billing_key_hash": res_billing_hash[:8] + "...", + "reserved_msats": res_msats, + }, + ) + continue + # Commit each release on its own so a later corrupt record's rollback + # cannot discard the healthy releases already processed in this batch. + await session.commit() released += 1 # Rolling upgrades can leave aggregate reservations created before durable @@ -246,16 +300,13 @@ async def release_stale_reservations( col(ApiKey.reserved_at) < cutoff ) else: - legacy_query = legacy_query.where( - or_( - col(ApiKey.hashed_key) == key_hash, - col(ApiKey.parent_key_hash) == key_hash, - ) - ).where( + legacy_query = legacy_query.where(col(ApiKey.hashed_key) == key_hash).where( or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff) ) for legacy_key in (await session.exec(legacy_query)).all(): + observed_reserved = legacy_key.reserved_balance + observed_reserved_at = legacy_key.reserved_at active_owner = ( await session.exec( select(ReservationRelease.id) @@ -272,10 +323,10 @@ async def release_stale_reservations( ).first() if active_owner is not None: continue - legacy_key.reserved_balance = 0 - legacy_key.reserved_at = None - session.add(legacy_key) - released += 1 + if await _release_legacy_aggregate( + session, legacy_key.hashed_key, observed_reserved, observed_reserved_at + ): + released += 1 await session.commit() if released: @@ -290,21 +341,15 @@ async def release_stale_reservations( async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> int: - """Delete dead parentless API keys; return the count removed. + """Delete dead API keys; return the count removed. - Dead = 0 balance/reservation/spend/requests, older than the grace period, - no parent, no children, no invoice that could still settle. Cashu rows are + Dead = 0 balance/reservation/spend/requests, older than the grace + period, no invoice that could still settle. Cashu rows are unlinked (not deleted) first to keep the audit trail. """ now = int(time.time()) cutoff = now - min_age_seconds - child = aliased(ApiKey) - has_children = ( - select(child.hashed_key).where( - col(child.parent_key_hash) == col(ApiKey.hashed_key) - ) - ).exists() # An expired invoice stays creditable for the grace window, and crediting it # after its target key is gone strands the payment at the mint. settleable_invoice = ( @@ -322,16 +367,22 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in ) ).exists() + has_refund_claim = ( + select(Refund.id).where( + col(Refund.api_key_hashed_key) == col(ApiKey.hashed_key) + ) + ).exists() + eligible_hashes = ( select(ApiKey.hashed_key) .where(col(ApiKey.balance) == 0) .where(col(ApiKey.reserved_balance) == 0) .where(col(ApiKey.total_spent) == 0) .where(col(ApiKey.total_requests) == 0) - .where(col(ApiKey.parent_key_hash).is_(None)) .where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff)) .where(~settleable_invoice) - .where(~has_children) + # refunds holds a non-null FK to the key. + .where(~has_refund_claim) ) # Unlink transactions rather than cascade-deleting them, so the financial @@ -386,9 +437,10 @@ class ModelRow(SQLModel, table=True): # type: ignore class ModelPathRow(SQLModel, table=True): # type: ignore """Upstream provider path a model is reachable through. - Discovery/visibility data only. ``model_id`` is intentionally NOT globally - unique: it is the client-visible ``/v1/models`` id (``forwarded_model_id or - id``) grouped across every provider that exposes the model. A single model + Discovery data plus provider-specific model metadata. ``model_id`` is + intentionally NOT globally unique: it is the client-visible ``/v1/models`` + id (``forwarded_model_id or id``) grouped across every provider that exposes + the model. A single model can therefore have several rows — one per direct provider path plus one per OpenRouter sub-provider endpoint. """ @@ -425,6 +477,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore endpoint_name: str | None = Field( default=None, description="Human-readable endpoint display name" ) + model_metadata: str = Field( + default="{}", + description="JSON model metadata specific to this provider path", + ) upstream_provider_id: int = Field( index=True, foreign_key="upstream_providers.id", @@ -470,14 +526,6 @@ class LightningInvoice(SQLModel, table=True): # type: ignore ) expires_at: int = Field(description="Unix timestamp when invoice expires") paid_at: int | None = Field(default=None, description="Unix timestamp when paid") - balance_limit: int | None = Field( - default=None, - description="Max spendable msats for the created key", - ) - balance_limit_reset: str | None = Field( - default=None, - description="Reset policy for balance limit (daily, weekly, monthly)", - ) validity_date: int | None = Field( default=None, description="Unix timestamp after which the created key expires", @@ -520,6 +568,53 @@ class CashuTransaction(SQLModel, table=True): # type: ignore ) +REFUND_OPEN_STATUSES = ("pending", "ambiguous") + +# Debited from the key but neither paid out nor restored, so still owed. +REFUND_UNRESOLVED_STATUSES = ("pending", "ambiguous", "stuck") + +_REFUND_OPEN_PREDICATE = "status IN ('pending', 'ambiguous')" + + +class Refund(SQLModel, table=True): # type: ignore + """One payout claim; the partial unique index allows one open claim per key.""" + + __tablename__ = "refunds" + __table_args__ = ( + Index( + "ux_refunds_open_per_key", + "api_key_hashed_key", + unique=True, + sqlite_where=text(_REFUND_OPEN_PREDICATE), + postgresql_where=text(_REFUND_OPEN_PREDICATE), + ), + ) + + id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex) + api_key_hashed_key: str = Field(foreign_key="api_keys.hashed_key", index=True) + method: str = Field(description="Payout method: lightning or cashu") + destination: str | None = Field( + default=None, description="Lightning address or LNURL, NULL for cashu" + ) + amount_msats: int = Field(description="Balance debited when the claim opened") + unit: str = Field(description="Mint unit the payout is denominated in") + mint_url: str = Field(description="Mint the payout is drawn from") + status: str = Field( + default="pending", + index=True, + description="pending, paid, failed, ambiguous, or stuck", + ) + quote_id: str | None = Field( + default=None, description="Melt quote id, for reconciling an ambiguous payout" + ) + token: str | None = Field(default=None, description="Issued cashu token") + claimed_at: int | None = Field( + default=None, description="Reconciler lease timestamp" + ) + created_at: int = Field(default_factory=lambda: int(time.time())) + updated_at: int = Field(default_factory=lambda: int(time.time())) + + async def store_cashu_transaction( token: str, amount: int, @@ -911,8 +1006,18 @@ async def complete_routstr_fee_payout( async def total_user_liability(db_session: AsyncSession) -> int: - """Return all outstanding API-key balances in millisatoshis.""" - result = await db_session.exec(select(func.sum(ApiKey.balance))) + """Return all outstanding user funds in millisatoshis. + + Key balances and unresolved refunds are summed in one statement so a + claim opened between two reads cannot be missed by both. + """ + key_balances = select(func.coalesce(func.sum(ApiKey.balance), 0)).scalar_subquery() + unresolved_refunds = ( + select(func.coalesce(func.sum(Refund.amount_msats), 0)) + .where(col(Refund.status).in_(REFUND_UNRESOLVED_STATUSES)) + .scalar_subquery() + ) + result = await db_session.exec(select(key_balances + unresolved_refunds)) return int(result.one() or 0) diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index 9e4f2ae6..2fc34bb6 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -1,4 +1,8 @@ +import math + from fastapi import Request +from fastapi.encoders import jsonable_encoder +from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse from .logging import get_logger @@ -30,6 +34,26 @@ class UpstreamError(Exception): super().__init__(message) +class EhbpTimeoutError(UpstreamError): + """Raised when an EHBP upstream times out waiting for a response. + + Distinct from a generic :class:`UpstreamError` so callers can map the + failure to a ``504 Gateway Timeout`` with a stable ``UPSTREAM_TIMEOUT`` + code instead of a misleading ``500`` internal server error. + + ``details`` carries optional structured, redaction-safe context and is + forwarded to the client by ``create_upstream_error_response``. + """ + + def __init__(self, message: str, details: dict[str, object] | None = None): + super().__init__( + message, + status_code=504, + code="UPSTREAM_TIMEOUT", + details=details, + ) + + async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse: """Handle HTTP exceptions and include request ID in response.""" request_id = getattr(request.state, "request_id", "unknown") @@ -40,15 +64,25 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon path = request.url.path # 4xx is client behaviour; the uvicorn access log already records it. - # Only 5xx warrants a server-side warning/error log here. if status_code >= 500: - logger.error( + error_type = None + if isinstance(detail, dict): + error = detail.get("error") + if isinstance(error, dict): + error_type = error.get("type") + log = ( + logger.warning + if error_type in {"mint_unreachable", "mint_rate_limited"} + else logger.error + ) + log( f"HTTP {status_code} on {path}: {detail}", extra={ "request_id": request_id, "status_code": status_code, "detail": detail, "path": path, + "error_type": error_type, }, ) @@ -61,6 +95,45 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon return JSONResponse(status_code=status_code, content=content) +def json_compliant(value: object) -> object: + """Render non-finite floats as text so a reply carrying them can serialize. + + ``json`` parses the bare ``NaN``/``Infinity``/``-Infinity`` literals into + real floats, so a request body — and a stored row written from one — may + hold one anywhere. ``JSONResponse`` encodes with ``allow_nan=False`` and + raises on them, which would turn a reply that merely *quotes* the offending + value into a 500. + """ + if isinstance(value, float) and not math.isfinite(value): + return repr(value) + if isinstance(value, dict): + return {key: json_compliant(item) for key, item in value.items()} + if isinstance(value, (list, tuple)): + return [json_compliant(item) for item in value] + return value + + +async def validation_exception_handler( + request: Request, exc: Exception +) -> JSONResponse: + """Answer a request-validation failure with a 422 that always serializes. + + Pydantic echoes the rejected value back in each error's ``input`` field. A + non-finite float there breaks the encoder, so the 422 escapes as a 500 and + reports a client's bad rate as a server fault. + """ + request_id = getattr(request.state, "request_id", "unknown") + errors = exc.errors() if isinstance(exc, RequestValidationError) else [] + + return JSONResponse( + status_code=422, + content={ + "detail": json_compliant(jsonable_encoder(errors)), + "request_id": request_id, + }, + ) + + async def general_exception_handler(request: Request, exc: Exception) -> JSONResponse: """Handle general exceptions and include request ID in response.""" request_id = getattr(request.state, "request_id", "unknown") diff --git a/routstr/core/logging.py b/routstr/core/logging.py index 3886637c..0fba407b 100644 --- a/routstr/core/logging.py +++ b/routstr/core/logging.py @@ -16,7 +16,7 @@ DO NOT modify or remove these messages without updating the usage tracking logic - The 'token_cost', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics 3. "Max cost payment finalized" (INFO) - routstr/auth.py - - Used as the successful completion fallback when token usage is unavailable + - Used for explicit flat-price/MaxCostData settlements; missing usage alone must not create this charge - The 'charged_amount', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics 4. "Payment processed successfully" (INFO) - routstr/auth.py @@ -51,7 +51,7 @@ from pythonjsonlogger import jsonlogger from rich.console import Console from rich.logging import RichHandler -from .redaction import redact_obj, redact_org_ids +from .redaction import redact_field, redact_org_ids # Only use RichHandler when stdout is a real TTY. In non-TTY contexts # (docker logs, pipes, CI) Rich pads every line to width and wraps long @@ -100,8 +100,10 @@ class DailyRotatingFileHandler(logging.handlers.TimedRotatingFileHandler): self.baseFilename = new_filename self.current_date = new_date - # FIX ME: not sure if we need this - # self._cleanup_old_files() + # `backupCount` alone never prunes these files: the base filename moves + # with the date, so the inherited rollover finds no siblings to expire + # and every day of logged credentials is retained indefinitely. + self._cleanup_old_files() if not self.delay: self.stream = self._open() @@ -182,6 +184,18 @@ class RequestIdFilter(logging.Filter): return True +class ClientAppFilter(logging.Filter): + """Attach request-local app attribution to log records.""" + + def filter(self, record: logging.LogRecord) -> bool: + # Import here to avoid circular imports + from .middleware import UNKNOWN_CLIENT_APP, client_app_context + + client_app = client_app_context.get(None) + record.client_app = client_app if client_app else UNKNOWN_CLIENT_APP + return True + + # Standard ``LogRecord`` attributes that are never user-supplied ``extra`` # fields; skipped when redacting structured extras (``msg``/``message`` are # handled separately above). @@ -260,13 +274,11 @@ class SecurityFilter(logging.Filter): # Structured `extra={...}` fields are emitted by the JSON formatter # straight from the record dict and never pass through the message - # formatting above. Redact organization IDs from any string-valued - # extra so they cannot leak via structured logs. + # formatting above, so they need their own recursive pass. for attr, value in list(record.__dict__.items()): if attr in _NON_EXTRA_RECORD_ATTRS: continue - if isinstance(value, (str, dict, list, tuple)): - record.__dict__[attr] = redact_obj(value) + record.__dict__[attr] = redact_field(attr, value) except Exception: pass @@ -323,7 +335,7 @@ def setup_logging() -> None: "rich_tracebacks": True, "markup": True, "console": _console, - "filters": ["request_id_filter", "security_filter"], + "filters": ["request_id_filter", "client_app_filter", "security_filter"], } else: console_handler = { @@ -331,7 +343,7 @@ def setup_logging() -> None: "level": log_level, "formatter": "plain", "stream": "ext://sys.stdout", - "filters": ["request_id_filter", "security_filter"], + "filters": ["request_id_filter", "client_app_filter", "security_filter"], } LOGGING_CONFIG = { @@ -340,7 +352,7 @@ def setup_logging() -> None: "formatters": { "json": { "()": jsonlogger.JsonFormatter, - "format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s", + "format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s %(client_app)s", "datefmt": "%Y-%m-%d %H:%M:%S", }, "plain": { @@ -351,6 +363,7 @@ def setup_logging() -> None: "filters": { "version_filter": {"()": VersionFilter}, "request_id_filter": {"()": RequestIdFilter}, + "client_app_filter": {"()": ClientAppFilter}, "security_filter": {"()": SecurityFilter}, }, "handlers": { @@ -364,7 +377,12 @@ def setup_logging() -> None: "interval": 1, # Every 1 day "backupCount": 30, # Keep 30 days of logs "atTime": None, # Rotate at midnight (00:00) - "filters": ["version_filter", "request_id_filter", "security_filter"], + "filters": [ + "version_filter", + "request_id_filter", + "client_app_filter", + "security_filter", + ], }, }, "loggers": { diff --git a/routstr/core/main.py b/routstr/core/main.py index 5ca71c9a..5ce22d22 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -4,6 +4,7 @@ from pathlib import Path from typing import AsyncGenerator from fastapi import FastAPI +from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse, RedirectResponse from fastapi.staticfiles import StaticFiles @@ -13,10 +14,10 @@ from starlette.types import Scope from ..auth import ( periodic_dead_key_prune, - periodic_key_reset, periodic_stale_reservation_sweep, ) from ..balance import balance_router, deprecated_wallet_router +from ..cashu_compat import install_cashu_httpx_shim from ..lightning import ( lightning_router, periodic_invoice_watcher, @@ -31,13 +32,18 @@ from ..nostr.discovery import providers_router from ..payment.models import models_router, update_sats_pricing from ..payment.price import update_prices_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically +from ..refund import periodic_refund_reconcile from ..upstream.auto_topup import periodic_auto_topup from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing from ..upstream.litellm_routing import configure_litellm from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout from .admin import admin_router from .db import create_session, init_db, run_migrations -from .exceptions import general_exception_handler, http_exception_handler +from .exceptions import ( + general_exception_handler, + http_exception_handler, + validation_exception_handler, +) from .logging import get_logger, setup_logging from .middleware import LoggingMiddleware from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401 @@ -63,15 +69,21 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: models_refresh_task = None model_maps_refresh_task = None model_paths_refresh_task = None - key_reset_task = None stale_reservation_task = None dead_key_prune_task = None auto_topup_task = None refund_sweep_task = None + refund_reconcile_task = None routstr_fee_task = None invoice_watcher_task = None try: + # cashu 0.20.x passes the `proxies` kwarg httpx removed in 0.28. + # routstr.wallet and routstr.payment.lnurl also install this at import; + # repeating it here keeps startup correct for any future module that + # reaches cashu's mint client without going through those two. + install_cashu_httpx_shim() + # Apply litellm-wide settings (drop_params, chat-completions URL, # debug logging) before any upstream provider dispatches a request. configure_litellm() @@ -143,16 +155,18 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: refresh_model_paths_periodically(get_upstreams) ) payout_task = asyncio.create_task(periodic_payout()) - if global_settings.nsec: - nip91_task = asyncio.create_task(announce_provider()) + # Always started: the loop idles until an NSEC is configured and re-reads + # it every iteration, so a key saved (or cleared) through the admin UI + # takes effect without a restart. + nip91_task = asyncio.create_task(announce_provider()) analytics_task = asyncio.create_task(publish_usage_analytics()) if global_settings.providers_refresh_interval_seconds > 0: providers_task = asyncio.create_task(providers_cache_refresher()) - key_reset_task = asyncio.create_task(periodic_key_reset()) stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep()) dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune()) auto_topup_task = asyncio.create_task(periodic_auto_topup()) refund_sweep_task = asyncio.create_task(periodic_refund_sweep()) + refund_reconcile_task = asyncio.create_task(periodic_refund_reconcile()) routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout()) invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher()) @@ -188,8 +202,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: model_maps_refresh_task.cancel() if model_paths_refresh_task is not None: model_paths_refresh_task.cancel() - if key_reset_task is not None: - key_reset_task.cancel() if stale_reservation_task is not None: stale_reservation_task.cancel() if dead_key_prune_task is not None: @@ -198,6 +210,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: auto_topup_task.cancel() if refund_sweep_task is not None: refund_sweep_task.cancel() + if refund_reconcile_task is not None: + refund_reconcile_task.cancel() if routstr_fee_task is not None: routstr_fee_task.cancel() if invoice_watcher_task is not None: @@ -223,8 +237,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(model_maps_refresh_task) if model_paths_refresh_task is not None: tasks_to_wait.append(model_paths_refresh_task) - if key_reset_task is not None: - tasks_to_wait.append(key_reset_task) if stale_reservation_task is not None: tasks_to_wait.append(stale_reservation_task) if dead_key_prune_task is not None: @@ -233,6 +245,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(auto_topup_task) if refund_sweep_task is not None: tasks_to_wait.append(refund_sweep_task) + if refund_reconcile_task is not None: + tasks_to_wait.append(refund_reconcile_task) if routstr_fee_task is not None: tasks_to_wait.append(routstr_fee_task) if invoice_watcher_task is not None: @@ -293,6 +307,7 @@ app.add_middleware(LoggingMiddleware) # Add exception handlers app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore +app.add_exception_handler(RequestValidationError, validation_exception_handler) app.add_exception_handler(Exception, general_exception_handler) @@ -306,7 +321,6 @@ async def info() -> dict: "mints": global_settings.cashu_mints, "http_url": global_settings.http_url, "onion_url": global_settings.onion_url, - "child_key_cost_msats": global_settings.child_key_cost, } diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py index 442c0a18..d4ddfb8f 100644 --- a/routstr/core/middleware.py +++ b/routstr/core/middleware.py @@ -2,8 +2,10 @@ import time import uuid from contextvars import ContextVar from typing import Callable +from urllib.parse import urlsplit from fastapi import Request, Response +from starlette.datastructures import Headers from starlette.middleware.base import BaseHTTPMiddleware from .logging import get_logger @@ -13,6 +15,41 @@ logger = get_logger(__name__) # Context variable to store request ID across async context request_id_context: ContextVar[str | None] = ContextVar("request_id") +client_app_context: ContextVar[str | None] = ContextVar("client_app") + +UNKNOWN_CLIENT_APP = "unknown" + +# Prefer OpenRouter app headers, then browser and SDK fallbacks. +_CLIENT_APP_HEADERS: tuple[str, ...] = ( + "x-title", + "http-referer", + "referer", + "user-agent", +) + +# Limit untrusted header data repeated in every log record. +_CLIENT_APP_MAX_LENGTH = 120 + + +def client_app_from_headers(headers: Headers) -> str: + for header in _CLIENT_APP_HEADERS: + raw = headers.get(header) + if raw is None: + continue + cleaned = "".join(ch for ch in raw if ch.isprintable()).strip() + if header in ("http-referer", "referer"): + try: + url = urlsplit(cleaned) + if url.scheme not in ("http", "https") or not url.hostname: + continue + except ValueError: + continue + # Attribution needs the origin, not credentials or private page URLs. + cleaned = f"{url.scheme}://{url.netloc.rsplit('@', 1)[-1]}" + if cleaned: + return cleaned[:_CLIENT_APP_MAX_LENGTH] + return UNKNOWN_CLIENT_APP + # Methods that are never logged: HEAD requests are health probes from # monitoring/load balancers, OPTIONS are CORS preflights — both are framework @@ -71,6 +108,10 @@ class LoggingMiddleware(BaseHTTPMiddleware): # Set request ID in context for logging token = request_id_context.set(request_id) + client_app_token = client_app_context.set( + client_app_from_headers(request.headers) + ) + path = request.url.path should_log = _should_log(request.method, path) @@ -84,7 +125,9 @@ class LoggingMiddleware(BaseHTTPMiddleware): "request_id": request_id, "method": request.method, "path": path, - "query_params": dict(request.query_params), + # Names only: query values carry API keys and refund + # tokens on the wallet routes. + "query_param_names": sorted(request.query_params.keys()), }, ) @@ -128,6 +171,12 @@ class LoggingMiddleware(BaseHTTPMiddleware): finally: # Reset context request_id_context.reset(token) + client_app_context.reset(client_app_token) -__all__ = ["LoggingMiddleware", "request_id_context"] +__all__ = [ + "LoggingMiddleware", + "UNKNOWN_CLIENT_APP", + "client_app_context", + "request_id_context", +] diff --git a/routstr/core/redaction.py b/routstr/core/redaction.py index c0519581..bf8d9cc6 100644 --- a/routstr/core/redaction.py +++ b/routstr/core/redaction.py @@ -1,8 +1,9 @@ -"""Redaction helpers for sensitive provider identifiers. +"""Redaction helpers for sensitive provider identifiers and credentials. Single source of truth for stripping account-scoped identifiers (e.g. OpenAI -organization IDs) from any text before it is logged, returned to a caller, or -written to an audit entry. +organization IDs) and spendable credentials (Cashu tokens, bearer keys, key +hashes) from any text before it is logged, returned to a caller, or written to +an audit entry. """ from __future__ import annotations @@ -33,19 +34,101 @@ def redact_org_ids(text: str) -> str: return _ORG_ID_PATTERN.sub(ORG_ID_PLACEHOLDER, text) -def redact_obj(obj: Any) -> Any: - """Recursively redact organization IDs in arbitrary nested structures. +SECRET_PLACEHOLDER = "[REDACTED]" - Strings are redacted in place; dicts and lists/tuples are walked so that - identifiers nested inside structured payloads (e.g. log ``extra`` fields or - error ``details``) are also stripped. Other types are returned unchanged. - """ +# Field names whose value is spendable or authenticating on its own. Matched as +# substrings of the lowercased key, so ``hashed_key`` (a live ``sk-`` credential) +# is stripped while the truncated ``key_hash`` prefix used for correlation is +# not. Numeric values are never stripped, which keeps ``input_tokens`` and the +# other usage-analytics fields intact. +_SECRET_KEY_HINTS = ( + "authorization", + "api_key", + "apikey", + "bearer", + "cashu", + "cookie", + "credential", + "hashed_key", + "mnemonic", + "nsec", + "passphrase", + "password", + "private_key", + "privkey", + "secret", + "token", +) + +# Value shapes that are spendable wherever they appear, including inside URLs, +# query strings and free-form error text. Every alternative is anchored on a +# literal prefix and uses a single bounded character class, so matching stays +# linear on the logging hot path. +_SECRET_VALUE_PATTERNS: tuple[re.Pattern[str], ...] = ( + re.compile(r"cashu[A-Z][A-Za-z0-9_\-=/+]{20,}"), + re.compile(r"\bnsec1[a-z0-9]{20,}"), + re.compile(r"\bBearer\s+[A-Za-z0-9_\-.=]{10,}", re.IGNORECASE), + re.compile(r"\bsk-[A-Za-z0-9]{16,}"), + re.compile(r"\b[0-9a-f]{64}\b"), + re.compile( + r"(?<=[?&])([^=&\s]*(?:token|key|secret|password|auth|sig)[^=&\s]*=)[^&\s\"']+", + re.IGNORECASE, + ), +) + +_MAX_REDACTION_DEPTH = 12 + + +def _redact_secret_text(text: str) -> str: + for pattern in _SECRET_VALUE_PATTERNS: + text = pattern.sub( + lambda match: (match.group(1) if match.groups() else "") + + SECRET_PLACEHOLDER, + text, + ) + return redact_org_ids(text) + + +def _is_secret_key(key: object) -> bool: + if not isinstance(key, str): + return False + lowered = key.lower() + return any(hint in lowered for hint in _SECRET_KEY_HINTS) + + +def _redact_secrets(obj: Any, depth: int, seen: frozenset[int]) -> Any: if isinstance(obj, str): - return redact_org_ids(obj) + return _redact_secret_text(obj) + if not isinstance(obj, (dict, list, tuple)): + return obj + if depth >= _MAX_REDACTION_DEPTH or id(obj) in seen: + return SECRET_PLACEHOLDER + nested = seen | {id(obj)} if isinstance(obj, dict): - return {key: redact_obj(value) for key, value in obj.items()} - if isinstance(obj, list): - return [redact_obj(value) for value in obj] - if isinstance(obj, tuple): - return tuple(redact_obj(value) for value in obj) - return obj + return { + key: _redact_value(key, value, depth + 1, nested) + for key, value in obj.items() + } + redacted = [_redact_secrets(value, depth + 1, nested) for value in obj] + return redacted if isinstance(obj, list) else tuple(redacted) + + +def _redact_value(key: object, value: Any, depth: int, seen: frozenset[int]) -> Any: + # Containers keep being walked even under a secret-shaped key so that the + # surrounding structure stays readable for operators. + if isinstance(value, (bool, int, float, dict, list, tuple)) or value is None: + return _redact_secrets(value, depth, seen) + if _is_secret_key(key): + return SECRET_PLACEHOLDER + return _redact_secrets(value, depth, seen) + + +def redact_field(key: str, value: Any) -> Any: + """Strip credentials from one named field, e.g. a log ``extra`` entry. + + Both secret-shaped keys and secret-shaped values are stripped, so a leak + survives neither a renamed field nor a credential embedded in free text. + The walk is depth-limited and cycle-safe: a malformed payload degrades to + ``[REDACTED]`` rather than taking the logging call down with it. + """ + return _redact_value(key, value, 0, frozenset()) diff --git a/routstr/core/settings.py b/routstr/core/settings.py index c02c6030..2104aeb3 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -77,7 +77,6 @@ class Settings(BaseSettings): exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE") upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE") tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE") - child_key_cost: int = Field(default=0, env="CHILD_KEY_COST") # Minimum per-request charge in millisatoshis when model pricing is free/zero min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT") reset_reserved_balance_on_startup: bool = Field( @@ -101,6 +100,12 @@ class Settings(BaseSettings): # Network cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") + # Comma-separated METHOD:path pairs adding to the proxy's canonical + # endpoint allowlist, e.g. "POST:v1/rerank,GET:batches". Only for upstreams + # exposing an endpoint outside the OpenAI-compatible set; each addition + # widens what the provider credential can be spent against, so wildcards + # and prefixes are not supported. + proxy_extra_allowed_paths: str = Field(default="", env="PROXY_EXTRA_ALLOWED_PATHS") tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL") providers_refresh_interval_seconds: int = Field( default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" @@ -119,7 +124,6 @@ class Settings(BaseSettings): enable_model_paths_refresh: bool = Field( default=True, env="ENABLE_MODEL_PATHS_REFRESH" ) - refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") # Uncollected refund tokens are swept after ~6 months (180 days). # Fixed for now: not configurable via env or the settings DB/admin API # (empty env list disables env binding; see FIXED_FIELDS). @@ -127,6 +131,14 @@ class Settings(BaseSettings): refund_sweep_claim_timeout_seconds: int = Field( default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS" ) + # How long an open refund claim may sit before the reconciler asks the mint + # what became of it. Doubles as the reconciler's per-row lease. + refund_claim_timeout_seconds: int = Field( + default=300, gt=0, env="REFUND_CLAIM_TIMEOUT_SECONDS" + ) + refund_reconcile_interval_seconds: int = Field( + default=60, gt=0, env="REFUND_RECONCILE_INTERVAL_SECONDS" + ) # Database connection-pool controls (advanced). Capacity defaults provide # headroom for Routstr's concurrent request and background-payment workload. @@ -142,6 +154,9 @@ class Settings(BaseSettings): database_pool_hold_warn_seconds: float = Field( default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS" ) + database_busy_timeout: float = Field( + default=30.0, gt=0, env="DATABASE_BUSY_TIMEOUT" + ) # Logging log_level: str = Field(default="INFO", env="LOG_LEVEL") @@ -192,12 +207,18 @@ SECRET_FIELDS = frozenset({"admin_password", "nsec"}) # neither store nor shadow them; env is always authoritative. ENV_ONLY_FIELDS = frozenset( { + # Widening the proxy's reachable upstream surface is a deployment + # decision, not a runtime toggle: it changes what the provider + # credential can be spent against. Keeping it env-only also lets the + # proxy parse it once at import without going stale. + "proxy_extra_allowed_paths", "database_pool_size", "database_max_overflow", "database_pool_timeout", "database_pool_recycle", "database_pool_pre_ping", "database_pool_hold_warn_seconds", + "database_busy_timeout", } ) @@ -248,7 +269,7 @@ def derive_npub_from_nsec(nsec: str) -> str | None: boot. """ try: - from nostr.key import PublicKey # type: ignore + from nostr_sdk import PublicKey from ..nostr.listing import nsec_to_keypair except ImportError: @@ -260,7 +281,7 @@ def derive_npub_from_nsec(nsec: str) -> str | None: _privkey_hex, pubkey_hex = keypair try: - return PublicKey(bytes.fromhex(pubkey_hex)).bech32() + return PublicKey.parse(pubkey_hex).to_bech32() except (ValueError, AttributeError): return None diff --git a/routstr/core/version.py b/routstr/core/version.py index b1633047..b2381ae3 100644 --- a/routstr/core/version.py +++ b/routstr/core/version.py @@ -20,7 +20,7 @@ import subprocess from functools import lru_cache from pathlib import Path -BASE_VERSION = "0.4.5" +BASE_VERSION = "0.4.7" _REPO_ROOT = Path(__file__).resolve().parents[2] _GIT_TIMEOUT_SECONDS = 2.0 diff --git a/routstr/lightning.py b/routstr/lightning.py index e7b8046a..b3ece68f 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -80,8 +80,6 @@ class _InvoiceSettlement: purpose: str api_key_hash: str | None mint_url: str | None - balance_limit: int | None - balance_limit_reset: str | None validity_date: int | None @classmethod @@ -93,8 +91,6 @@ class _InvoiceSettlement: purpose=invoice.purpose, api_key_hash=invoice.api_key_hash, mint_url=invoice.mint_url, - balance_limit=invoice.balance_limit, - balance_limit_reset=invoice.balance_limit_reset, validity_date=invoice.validity_date, ) @@ -118,8 +114,6 @@ class InvoiceCreateRequest(BaseModel): default=None, description="Deprecated: legacy field for topup. Prefer Authorization header.", ) - balance_limit: int | None = Field(default=None) - balance_limit_reset: str | None = Field(default=None) validity_date: int | None = Field(default=None) @@ -220,11 +214,18 @@ async def _request_mint_with_fallback( ) continue try: - wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False) + wallet = await get_wallet( + mint_url, + "sat", + retry_on_rate_limit=False, + load_proofs=False, + ) quote = await run_mint_operation( lambda: wallet.request_mint(amount_sats), op_name="request_mint_invoice", mint_url=mint_url, + # Quote creation is unsafe to retry without idempotency. + retry_timeouts=False, retry_on_rate_limit=False, ) return quote.request, quote.quote, mint_url @@ -400,8 +401,6 @@ async def create_invoice( api_key_hash=api_key_token[3:] if api_key_token else None, purpose=request.purpose, mint_url=mint_url, - balance_limit=request.balance_limit, - balance_limit_reset=request.balance_limit_reset, validity_date=request.validity_date, expires_at=expires_at, ) @@ -582,7 +581,7 @@ async def check_invoice_payment( await session.commit() mint_url = settlement.mint_url or settings.primary_mint - wallet = await get_wallet(mint_url, "sat") + wallet = await get_wallet(mint_url, "sat", load_proofs=False) try: mint_status = await run_mint_operation( lambda: wallet.get_mint_quote(settlement.payment_hash), @@ -797,8 +796,6 @@ async def _create_api_key_record( balance=invoice.amount_sats * 1000, refund_currency="sat", refund_mint_url=mint_url, - balance_limit=invoice.balance_limit, - balance_limit_reset=invoice.balance_limit_reset, validity_date=invoice.validity_date, ) session.add(api_key) diff --git a/routstr/mint.py b/routstr/mint.py index 64632b9e..e1b60ca2 100644 --- a/routstr/mint.py +++ b/routstr/mint.py @@ -177,6 +177,8 @@ class MintRateGuard: if isinstance(error, httpx.HTTPStatusError): retry_after = parse_retry_after(error.response.headers) self.apply_rate_limit_cooldown(retry_after) + elif is_mint_transport_error(error): + self.apply_cooldown(MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport") else: self.apply_cooldown(1.0) logger.warning( @@ -236,6 +238,17 @@ def mint_cooldown_reason(mint_url: str) -> str | None: return MintRateGuard.get(mint_url).cooldown_reason() +def is_mint_transport_error(error: BaseException) -> bool: + current: BaseException | None = error + seen: set[int] = set() + while current is not None and id(current) not in seen: + seen.add(id(current)) + if isinstance(current, MINT_TRANSPORT_EXCEPTIONS): + return True + current = current.__cause__ or current.__context__ + return False + + def is_mint_rate_limited(error: BaseException) -> bool: """Return whether an exception chain represents HTTP 429/cooldown.""" @@ -269,6 +282,7 @@ async def run_mint_operation( mint_url: str = "", retry_timeouts: bool = True, retry_on_rate_limit: bool = True, + allow_during_cooldown: bool = False, ) -> Any: """Run one mint operation with bounded concurrency and adaptive cooldown.""" @@ -282,7 +296,7 @@ async def run_mint_operation( return await factory() async def invoke() -> Any: - if guard is not None: + if guard is not None and not allow_during_cooldown: return await guard.run(timed_factory) return await timed_factory() @@ -293,6 +307,7 @@ async def run_mint_operation( raise except (asyncio.TimeoutError, httpx.TimeoutException) as exc: if retry_timeouts and attempt < max_attempts - 1: + # Cooldown opens only after retries; earlier would stretch each backoff to a full cooldown wait. backoff = (2**attempt) + (time.monotonic() % 1.0) logger.warning( "Mint operation timed out, retrying", @@ -305,11 +320,19 @@ async def run_mint_operation( ) await asyncio.sleep(backoff) continue + if guard is not None: + guard.apply_cooldown( + MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport" + ) raise httpx.TimeoutException( f"{op_name} timed out (attempts: {attempt + 1})" ) from exc except Exception as exc: if not is_mint_rate_limited(exc): + if guard is not None and is_mint_transport_error(exc): + guard.apply_cooldown( + MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport" + ) raise backoff = (2**attempt) + (time.monotonic() % 1.0) diff --git a/routstr/nostr/analytics.py b/routstr/nostr/analytics.py index e568b5e0..54b7ecf4 100644 --- a/routstr/nostr/analytics.py +++ b/routstr/nostr/analytics.py @@ -12,13 +12,11 @@ import json import time from typing import Any -from nostr.event import Event -from nostr.key import PrivateKey - from ..core import get_logger from ..core.log_manager import log_manager from ..core.settings import settings from .listing import nsec_to_keypair, publish_to_relay +from .sdk import create_signed_event logger = get_logger(__name__) @@ -44,18 +42,6 @@ WINDOW_DEFINITIONS: tuple[tuple[str, int, int], ...] = ( ) -def _event_to_dict(ev: Event) -> dict[str, Any]: - return { - "id": ev.id, - "pubkey": ev.public_key, - "created_at": ev.created_at, - "kind": int(ev.kind) if not isinstance(ev.kind, int) else ev.kind, - "tags": ev.tags, - "content": ev.content, - "sig": ev.signature, - } - - def _resolve_provider_id(public_key_hex: str) -> str: explicit_provider_id = (settings.provider_id or "").strip() if explicit_provider_id: @@ -293,21 +279,18 @@ def create_stats_snapshot_event( *, d_tag: str, ) -> dict[str, Any]: - private_key = PrivateKey(bytes.fromhex(private_key_hex)) tags = [ ["d", d_tag], ["provider", provider_id], ["schema", ANALYTICS_SCHEMA], ] - event = Event( - public_key=private_key.public_key.hex(), - content=payload_json, + return create_signed_event( + private_key_hex, kind=ANALYTICS_KIND, + content=payload_json, tags=tags, ) - private_key.sign_event(event) - return _event_to_dict(event) def _fingerprint_payload(payload: dict[str, Any]) -> str: diff --git a/routstr/nostr/discovery.py b/routstr/nostr/discovery.py index 268b4c7c..ce61d4e3 100644 --- a/routstr/nostr/discovery.py +++ b/routstr/nostr/discovery.py @@ -320,18 +320,18 @@ async def fetch_provider_health(endpoint_url: str) -> dict[str, Any]: is_onion = ".onion" in endpoint_url # Set up client arguments conditionally - proxies = None + proxy: str | None = None if is_onion: try: tor_proxy = settings.tor_proxy_url except Exception: tor_proxy = "socks5://127.0.0.1:9050" - proxies = {"http://": tor_proxy, "https://": tor_proxy} # type: ignore[assignment] + proxy = tor_proxy async with httpx.AsyncClient( timeout=httpx.Timeout(30.0), follow_redirects=True, - proxies=proxies, # type: ignore[arg-type] + proxy=proxy, ) as client: # Prefer provider's /v1/info for full details info_url = f"{endpoint_url.rstrip('/')}/v1/info" diff --git a/routstr/nostr/listing.py b/routstr/nostr/listing.py index 7c60e9b1..02a668eb 100644 --- a/routstr/nostr/listing.py +++ b/routstr/nostr/listing.py @@ -8,18 +8,12 @@ import asyncio import json import os import random -import ssl import time from typing import Any, cast -from nostr.event import Event -from nostr.filter import Filter, Filters -from nostr.key import PrivateKey -from nostr.message_type import ClientMessageType -from nostr.relay_manager import RelayManager - from ..core import get_logger from ..core.settings import settings +from .sdk import create_signed_event, fetch_events, parse_keypair, send_event logger = get_logger(__name__) @@ -33,18 +27,6 @@ def get_app_version() -> str | None: return None -def _event_to_dict(ev: Event) -> dict[str, Any]: - return { - "id": ev.id, - "pubkey": ev.public_key, - "created_at": ev.created_at, - "kind": int(ev.kind) if not isinstance(ev.kind, int) else ev.kind, - "tags": ev.tags, - "content": ev.content, - "sig": ev.signature, - } - - def nsec_to_keypair(nsec: str) -> tuple[str, str] | None: """ Convert a Nostr private key (nsec) to a keypair (privkey_hex, pubkey_hex). @@ -56,16 +38,10 @@ def nsec_to_keypair(nsec: str) -> tuple[str, str] | None: Tuple of (private_key_hex, public_key_hex) or None if invalid """ try: - if nsec.startswith("nsec"): - pk = PrivateKey.from_nsec(nsec) - return (pk.hex(), pk.public_key.hex()) - - if len(nsec) == 64: - pk = PrivateKey(bytes.fromhex(nsec)) - return (pk.hex(), pk.public_key.hex()) - - logger.error(f"Invalid private key format/length: {len(nsec)}") - return None + if not (nsec.startswith("nsec") or len(nsec) == 64): + logger.error(f"Invalid private key format/length: {len(nsec)}") + return None + return parse_keypair(nsec) except Exception as e: logger.error(f"Failed to convert nsec to keypair: {e}") return None @@ -93,8 +69,6 @@ def create_listing_event( Returns: Complete signed nostr event as a dict ready for publishing """ - pk = PrivateKey(bytes.fromhex(private_key_hex)) - tags = [["d", provider_id]] for url in endpoint_urls: tags.append(["u", url]) @@ -107,9 +81,12 @@ def create_listing_event( content = json.dumps(metadata, separators=(",", ":")) if metadata else "" - ev = Event(pk.public_key.hex(), content, kind=38421, tags=tags) - pk.sign_event(ev) - return _event_to_dict(ev) + return create_signed_event( + private_key_hex, + kind=38421, + content=content, + tags=tags, + ) def _get_tag_values(event: dict[str, Any], key: str) -> list[str]: @@ -177,75 +154,25 @@ async def query_listing_events( succeeded without transport-level errors. """ - def _sync_query() -> tuple[list[dict[str, Any]], bool]: - rm = RelayManager() - rm.add_relay(relay_url) - events_out: list[dict[str, Any]] = [] - ok = True - try: - rm.open_connections({"cert_reqs": ssl.CERT_NONE}) - time.sleep(1.0) + try: + events_out = await fetch_events( + relay_url, + kind=38421, + author=pubkey, + limit=10, + timeout=timeout, + ) + except Exception as e: + logger.debug(f"Failed to query relay {relay_url}: {type(e).__name__}") + return [], False - flt = Filter(kinds=[38421], authors=[pubkey], limit=10) - filters = Filters([flt]) - sub_id = f"routstr_listing_{int(time.time())}" - rm.add_subscription(sub_id, filters) - req: list[Any] = [ClientMessageType.REQUEST, sub_id] - req.extend(filters.to_json_array()) - rm.publish_message(json.dumps(req)) - - start = time.time() - last_event_ts = start - while time.time() - start < timeout: - drained = False - while rm.message_pool.has_events(): - drained = True - ev_msg = rm.message_pool.get_event() - ev = ev_msg.event - ev_dict = _event_to_dict(ev) - if provider_id is not None: - tags = ev_dict.get("tags", []) - if not any( - isinstance(t, list) - and len(t) >= 2 - and t[0] == "d" - and t[1] == provider_id - for t in tags - ): - continue - events_out.append(ev_dict) - logger.debug( - f"Found listing event: {ev_dict.get('id', '')[:6]}...{ev_dict.get('id', '')[-6:]}" - ) - if drained: - last_event_ts = time.time() - - while rm.message_pool.has_notices(): - notice = rm.message_pool.get_notice() - try: - content = getattr(notice, "content", notice) - s = str(content) - if len(s) > 200: - s = s[:200] + "..." - logger.debug(f"Relay notice: {s}") - except Exception: - pass - - if time.time() - last_event_ts > 2.5: - break - - time.sleep(0.1) - except Exception as e: - ok = False - logger.debug(f"Failed to query relay {relay_url}: {type(e).__name__}") - finally: - try: - rm.close_connections() - except Exception: - pass - return events_out, ok - - return await asyncio.to_thread(_sync_query) + if provider_id is not None: + events_out = [ + event + for event in events_out + if _get_single_tag_value(event, "d") == provider_id + ] + return events_out, True def discover_onion_url_from_tor(base_dir: str = "/var/lib/tor") -> str | None: @@ -333,117 +260,102 @@ async def publish_to_relay( Publish a listing event to a nostr relay via nostr library. """ - def _sync_publish() -> bool: - rm = RelayManager() - rm.add_relay(relay_url) - try: - rm.open_connections({"cert_reqs": ssl.CERT_NONE}) - time.sleep(1.0) - # Publish the event as-is via publish_message to preserve signature - rm.publish_message(json.dumps(["EVENT", event])) - logger.debug(f"Sent listing event {event.get('id', '')} to {relay_url}") - time.sleep(1.0) - return True - except Exception as e: - logger.debug(f"Failed to publish to {relay_url}: {type(e).__name__}") - return False - finally: - try: - rm.close_connections() - except Exception: - pass - - return await asyncio.to_thread(_sync_publish) - - -async def announce_provider() -> None: - """ - Background task to announce this Routstr provider to Nostr relays. - Checks for existing announcements and creates new ones if needed. - """ - # Check for NSEC in environment (use NSEC only) - nsec = settings.nsec - if not nsec: - logger.info("Nostr private key not found (NSEC), skipping listing announcement") - return - - # Convert NSEC to keypair - keypair = nsec_to_keypair(nsec) - if not keypair: - logger.error("Failed to parse NSEC, skipping listing announcement") - return - - private_key_hex, public_key_hex = keypair - logger.info(f"Using Nostr pubkey: {public_key_hex}") - - # Resolve settings and determine if we can publish BEFORE touching relays try: - base_url: str | None = settings.http_url - onion_url: str | None = settings.onion_url - provider_name = settings.name or "Routstr Proxy" - provider_about = settings.description or "Privacy-preserving AI proxy via Nostr" - cashu_mints = [m.strip() for m in settings.cashu_mints if m.strip()] - except Exception: - base_url = settings.http_url or None - onion_url = settings.onion_url or None - provider_name = settings.name or "Routstr Proxy" - provider_about = settings.description or "Privacy-preserving AI proxy via Nostr" - cashu_mints = [m.strip() for m in settings.cashu_mints if m.strip()] + await send_event(relay_url, event, timeout=timeout) + logger.debug(f"Sent listing event {event.get('id', '')} to {relay_url}") + return True + except Exception as e: + logger.debug(f"Failed to publish to {relay_url}: {type(e).__name__}") + return False + + +# Re-announce cadence once a provider is listed. +ANNOUNCEMENT_INTERVAL_SECONDS = 24 * 60 * 60 +# Poll cadence while there is nothing to announce (no NSEC, no endpoint, ...). +DISABLED_POLL_SECONDS = 60 +# How often the long re-announce sleep re-checks the configured NSEC, so a +# newly saved identity is announced promptly instead of up to 24h later. +IDENTITY_POLL_SECONDS = 30 + +DEFAULT_RELAY_URLS = [ + "wss://relay.nostr.band", + "wss://relay.damus.io", + "wss://relay.routstr.com", + "wss://nos.lol", +] + + +def _resolve_endpoint_urls() -> list[str]: + """Endpoints to advertise: a public HTTP URL and/or an onion URL.""" + endpoint_urls: list[str] = [] + + base_url = (settings.http_url or "").strip() + if base_url and base_url != "http://localhost:8000": + endpoint_urls.append(base_url) + + onion_url = (settings.onion_url or "").strip() if not onion_url: discovered = discover_onion_url_from_tor() if discovered: onion_url = discovered logger.info(f"Discovered onion URL via Tor volume: {onion_url}") - mint_urls = cashu_mints if cashu_mints else None - endpoint_urls: list[str] = [] - if base_url and base_url.strip() and base_url.strip() != "http://localhost:8000": - endpoint_urls.append(base_url.strip()) - if onion_url and onion_url.strip(): - ou = onion_url.strip() - if ou.endswith(".onion") and not ( - ou.startswith("http://") or ou.startswith("https://") + if onion_url: + if onion_url.endswith(".onion") and not ( + onion_url.startswith("http://") or onion_url.startswith("https://") ): - ou = f"http://{ou}" - endpoint_urls.append(ou) + onion_url = f"http://{onion_url}" + endpoint_urls.append(onion_url) - if not endpoint_urls: - logger.warning( - "No valid endpoints configured (HTTP_URL/ONION_URL). Skipping listing publish." - ) - return + return endpoint_urls - # Only now configure relays and determine provider_id (may query relays) + +def _resolve_relay_urls() -> list[str]: relay_urls = [u.strip() for u in getattr(settings, "relays", []) if u.strip()] - if not relay_urls: - relay_urls = [ - "wss://relay.nostr.band", - "wss://relay.damus.io", - "wss://relay.routstr.com", - "wss://nos.lol", - ] + return relay_urls or list(DEFAULT_RELAY_URLS) - provider_id = await _determine_provider_id(public_key_hex, relay_urls) - logger.info(f"Using provider_id: {provider_id}") - # Build metadata - metadata = { - "name": provider_name, - "about": provider_about, - } +def _resolve_mint_urls() -> list[str] | None: + mints = [m.strip() for m in (settings.cashu_mints or []) if m.strip()] + return mints or None - # Create the candidate event that we would publish - version_str = get_app_version() - candidate_event = create_listing_event( - private_key_hex=private_key_hex, - provider_id=provider_id, - endpoint_urls=endpoint_urls, - mint_urls=mint_urls, - version=version_str, - metadata=metadata, - ) - # Backoff configuration and state (sensible defaults) +async def _sleep_until_next_announcement( + seconds: float, parsed_nsec: str | None +) -> None: + """Sleep up to ``seconds``, returning early if the configured NSEC changes. + + Without the early wake, a node runner who replaces the NSEC in the admin UI + would wait out the whole re-announce interval before the new identity (and, + with it, the new ``d`` tag and npub) is announced. + """ + remaining = float(seconds) + while remaining > 0: + if (settings.nsec or "").strip() != (parsed_nsec or ""): + return + chunk = min(float(IDENTITY_POLL_SECONDS), remaining) + await asyncio.sleep(chunk) + remaining -= chunk + + +async def announce_provider() -> None: + """Background task announcing this Routstr provider to Nostr relays. + + Started unconditionally at boot: while the node has no NSEC the task idles + and re-checks, so an identity configured later through the admin UI is + picked up (and announced) without a restart. The identity, endpoints, mints + and relays are all re-read every iteration, mirroring + ``publish_usage_analytics``. + """ + parsed_nsec: str | None = None + private_key_hex: str | None = None + public_key_hex: str | None = None + provider_id: str | None = None + warned_missing_nsec = False + + # Backoff state is deliberately long-lived: it has to survive an idle poll, + # a full re-announce cycle and an identity change, otherwise a failing relay + # would be retried at full rate on every pass. backoff_base = 5.0 backoff_max = 900.0 backoff_jitter_ratio = 0.2 @@ -468,67 +380,74 @@ async def announce_provider() -> None: f"Backoff: {relay} delay={delay:.1f}s jitter={jitter:.1f}s next={int(scheduled)}" ) - # Fetch existing events for this provider_id - existing_events: list[dict[str, Any]] = [] - for relay_url in relay_urls: - if _should_skip(relay_url): - logger.debug(f"Skipping {relay_url} due to backoff") - continue - events, ok = await query_listing_events(relay_url, public_key_hex, provider_id) - if ok: - _register_success(relay_url) - existing_events.extend(events) - else: - _register_failure(relay_url) - - # Decide whether to publish: publish if none exist or any differ from candidate - found_any = len(existing_events) > 0 - all_match = found_any and all( - events_semantically_equal(ev, candidate_event) for ev in existing_events - ) - - if not all_match: - logger.debug( - "No matching listing announcement found or differences detected; publishing update" - ) - success_count = 0 - for relay_url in relay_urls: - if _should_skip(relay_url): - logger.debug(f"Skipping publish to {relay_url} due to backoff") - continue - if await publish_to_relay(relay_url, candidate_event): - _register_success(relay_url) - success_count += 1 - else: - _register_failure(relay_url) - logger.info( - f"Published listing announcement to {success_count}/{len(relay_urls)} relays" - ) - else: - logger.debug( - "Matching listing announcement already present; skipping publish on startup" - ) - - # Re-announce periodically (every 24 hours) - announcement_interval = 24 * 60 * 60 - while True: try: - await asyncio.sleep(announcement_interval) + nsec = (settings.nsec or "").strip() + + if not nsec: + if not warned_missing_nsec: + logger.info( + "Nostr private key not configured (NSEC); waiting for one " + "to be set before announcing this provider" + ) + warned_missing_nsec = True + parsed_nsec = None + await asyncio.sleep(DISABLED_POLL_SECONDS) + continue + + # Re-derive the identity whenever the configured NSEC changes, so a + # key saved (or replaced) through the admin UI takes effect live. + if nsec != parsed_nsec: + keypair = nsec_to_keypair(nsec) + if not keypair: + logger.error( + "Invalid NSEC; waiting for a valid one before announcing" + ) + parsed_nsec = None + await asyncio.sleep(DISABLED_POLL_SECONDS) + continue + private_key_hex, public_key_hex = keypair + parsed_nsec = nsec + provider_id = None + warned_missing_nsec = False + logger.info(f"Using Nostr pubkey: {public_key_hex}") + + if private_key_hex is None or public_key_hex is None: + await asyncio.sleep(DISABLED_POLL_SECONDS) + continue + + endpoint_urls = _resolve_endpoint_urls() + if not endpoint_urls: + logger.warning( + "No valid endpoints configured (HTTP_URL/ONION_URL). " + "Skipping listing publish until one is set." + ) + await asyncio.sleep(DISABLED_POLL_SECONDS) + continue + + relay_urls = _resolve_relay_urls() + + if provider_id is None: + provider_id = await _determine_provider_id(public_key_hex, relay_urls) + logger.info(f"Using provider_id: {provider_id}") + + metadata = { + "name": settings.name or "Routstr Proxy", + "about": settings.description + or "Privacy-preserving AI proxy via Nostr", + } - # Build fresh candidate event for comparison - version_str = get_app_version() candidate_event = create_listing_event( private_key_hex=private_key_hex, provider_id=provider_id, endpoint_urls=endpoint_urls, - mint_urls=mint_urls, - version=version_str, + mint_urls=_resolve_mint_urls(), + version=get_app_version(), metadata=metadata, ) # Fetch existing events for this provider_id - existing_events = [] + existing_events: list[dict[str, Any]] = [] for relay_url in relay_urls: if _should_skip(relay_url): logger.debug(f"Skipping {relay_url} due to backoff") @@ -549,26 +468,36 @@ async def announce_provider() -> None: if all_match: logger.debug( - "Matching listing announcement already present; skipping periodic re-announce" + "Matching listing announcement already present; skipping publish" + ) + else: + logger.debug( + "No matching listing announcement found or differences " + "detected; publishing update" + ) + success_count = 0 + for relay_url in relay_urls: + if _should_skip(relay_url): + logger.debug(f"Skipping publish to {relay_url} due to backoff") + continue + if await publish_to_relay(relay_url, candidate_event): + _register_success(relay_url) + success_count += 1 + else: + _register_failure(relay_url) + logger.info( + "Published listing announcement to " + f"{success_count}/{len(relay_urls)} relays" ) - continue - logger.debug( - f"Re-announcing provider due to differences or absence: {candidate_event['id']}" + # Re-announce periodically; wakes early if the NSEC changes. + await _sleep_until_next_announcement( + ANNOUNCEMENT_INTERVAL_SECONDS, parsed_nsec ) - for relay_url in relay_urls: - if _should_skip(relay_url): - logger.debug(f"Skipping publish to {relay_url} due to backoff") - continue - ok = await publish_to_relay(relay_url, candidate_event) - if ok: - _register_success(relay_url) - else: - _register_failure(relay_url) except asyncio.CancelledError: logger.info("Listing announcement task cancelled") break except Exception as e: logger.debug(f"Error in listing announcement loop: {type(e).__name__}") - # Continue running despite errors + await asyncio.sleep(DISABLED_POLL_SECONDS) diff --git a/routstr/nostr/sdk.py b/routstr/nostr/sdk.py new file mode 100644 index 00000000..e9e6a37f --- /dev/null +++ b/routstr/nostr/sdk.py @@ -0,0 +1,94 @@ +"""Small adapter around the maintained ``nostr-sdk`` package.""" + +from __future__ import annotations + +import json +from datetime import timedelta +from typing import Any, cast + +from nostr_sdk import ( + AckPolicy, + Client, + Event, + EventBuilder, + Filter, + Keys, + Kind, + PublicKey, + RelayUrl, + ReqExitPolicy, + ReqTarget, + SendEventTarget, + Tag, +) + + +def parse_keypair(secret_key: str) -> tuple[str, str]: + keys = Keys.parse(secret_key) + return keys.secret_key().to_hex(), keys.public_key().to_hex() + + +def create_signed_event( + secret_key_hex: str, + *, + kind: int, + content: str, + tags: list[list[str]], +) -> dict[str, Any]: + keys = Keys.parse(secret_key_hex) + event = ( + EventBuilder(Kind(kind), content) + .tags([Tag.parse(tag) for tag in tags]) + .finalize(keys) + ) + return cast(dict[str, Any], json.loads(event.as_json())) + + +async def fetch_events( + relay_url: str, + *, + kind: int, + author: str, + limit: int, + timeout: int, +) -> list[dict[str, Any]]: + relay = RelayUrl.parse(relay_url) + client = Client() + await client.add_relay(relay) + try: + await client.connect() + event_filter = ( + Filter().kinds([Kind(kind)]).authors([PublicKey.parse(author)]).limit(limit) + ) + events = await client.fetch_events( + ReqTarget.single(relay, [event_filter]), + timeout=timedelta(seconds=timeout), + policy=ReqExitPolicy.WAIT_DURATION_AFTER_EOSE(timedelta(seconds=2.5)), + max_events=limit, + ) + return [cast(dict[str, Any], json.loads(event.as_json())) for event in events] + finally: + await client.shutdown() + + +async def send_event(relay_url: str, event: dict[str, Any], *, timeout: int) -> None: + relay = RelayUrl.parse(relay_url) + client = Client() + await client.add_relay(relay) + try: + await client.connect() + output = await client.send_event( + Event.from_json(json.dumps(event)), + target=SendEventTarget.to([relay]), + ack_policy=AckPolicy.all(), + ok_timeout=timedelta(seconds=timeout), + ) + if output.failed or relay not in output.success: + reasons = ", ".join( + f"{url}: {reason}" for url, reason in output.failed.items() + ) + raise RuntimeError( + f"Relay did not accept event: {reasons or 'no OK from relay'}" + ) + finally: + await client.shutdown() diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 0fb305cb..b4b0ff6b 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -1,11 +1,12 @@ import math from typing import TYPE_CHECKING -from pydantic.v1 import BaseModel +from pydantic.v1 import BaseModel, Field from ..core import get_logger from ..core.settings import settings from .price import sats_usd_price +from .rates import coerce_rate, is_usable_rate from .usage import normalize_usage, parse_token_count if TYPE_CHECKING: @@ -34,6 +35,9 @@ class CostData(BaseModel): cache_creation_input_tokens: int = 0 cache_read_msats: int = 0 cache_creation_msats: int = 0 + # Actual debit after finalization; None means settlement has not run yet. + charged_msats: int | None = None + upstream_usd: float = Field(default=0.0, exclude=True) class MaxCostData(CostData): @@ -107,11 +111,11 @@ async def calculate_cost( if usage is None: logger.warning( - "No usage data in response — billing at MaxCostData with zero " - "tokens. Dashboard will show this request as `(0+0)`. Most " - "common cause: upstream stream did not include a final usage " - "chunk (OpenAI-compat backends require " - "`stream_options.include_usage=true`).", + "No usage data or local estimate in response — releasing the " + "reservation without charging it as usage. Dashboard will show " + "this request as `(0+0)` tokens. Most common cause: upstream " + "stream did not include a final usage chunk (OpenAI-compat " + "backends require `stream_options.include_usage=true`).", extra={ "max_cost_msats": max_cost, "model": response_data.get("model", "unknown"), @@ -175,13 +179,14 @@ async def calculate_cost( cost_details = usage_data.get("cost_details", {}) if not isinstance(cost_details, dict): cost_details = {} - input_usd = _coerce_usd( - cost_details.get("input_cost") - or cost_details.get("upstream_inference_prompt_cost") + # Coerce each spelling before choosing between them: `inf` and `NaN` + # are truthy, so a malformed first field would otherwise win the + # fallback and the usable figure beside it would never be read. + input_usd = _coerce_usd(cost_details.get("input_cost")) or _coerce_usd( + cost_details.get("upstream_inference_prompt_cost") ) - output_usd = _coerce_usd( - cost_details.get("output_cost") - or cost_details.get("upstream_inference_completions_cost") + output_usd = _coerce_usd(cost_details.get("output_cost")) or _coerce_usd( + cost_details.get("upstream_inference_completions_cost") ) cache_pricing_rates: tuple[float, float, float, float] | None = None if cache_read_tokens > 0 or cache_creation_tokens > 0: @@ -241,27 +246,32 @@ async def calculate_cost( else: input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates - if not (input_rate and output_rate): + # Truthiness is not the question: `NaN` and a negative rate are both truthy + # and sailed past this gate into the token math, while a rate of zero is a + # price — free — and reading it as a missing one charged the whole + # reservation for a request the model serves for nothing. + rates = (input_rate, output_rate, cache_read_rate, cache_creation_rate) + if not all(is_usable_rate(rate) for rate in rates): logger.warning( - "No token pricing configured — billing at flat MaxCostData. " - "Token counts %s in the upstream response but cannot be " - "priced; the request will appear in dashboards with the " - "raw counts and a fixed max-cost charge.", - "are present" - if (input_tokens > 0 or output_tokens > 0) - else "are zero", + "No usable token pricing — releasing the reservation instead of " + "treating its ceiling as the charge. Token counts %s in the " + "upstream response but cannot be converted to money; the request " + "will appear in dashboards with raw counts and a zero charge.", + "are present" if (input_tokens > 0 or output_tokens > 0) else "are zero", extra={ "base_cost_msats": max_cost, "model": response_data.get("model", "unknown"), "input_tokens": input_tokens, "output_tokens": output_tokens, + "input_rate": input_rate, + "output_rate": output_rate, }, ) return MaxCostData( - base_msats=max_cost, + base_msats=0, input_msats=0, output_msats=0, - total_msats=max_cost, + total_msats=0, input_tokens=input_tokens, output_tokens=output_tokens, cache_read_input_tokens=cache_read_tokens, @@ -289,15 +299,25 @@ async def calculate_cost( def _coerce_usd(value: object) -> float: - """Coerce a value to USD float, handling various formats safely.""" - if value is None or isinstance(value, bool): - return 0.0 - if not isinstance(value, (int, float, str)): - return 0.0 - try: - return max(0.0, float(value)) - except (TypeError, ValueError): - return 0.0 + """Coerce an upstream-reported USD figure to a usable amount, else ``0.0``. + + These values come straight off the upstream response, where ``json.loads`` + accepts the bare ``NaN``/``Infinity`` literals and overflows ``1e999`` to + ``inf``. A non-finite figure is not a cost, and letting one through poisoned + the proportional split in ``_calculate_from_usd_cost`` (``inf / inf`` is + ``NaN``): the resulting exception was absorbed by the broad handler around + the USD path, so a request whose *total* cost was perfectly valid fell + through to token-estimated pricing and was billed a fraction of what the + upstream charged. + + ``0.0`` means "no usable figure" to every caller, which is the same thing an + absent field means, so the caller's existing ``> 0`` checks handle it. + """ + # A cost figure is coerced exactly like a rate; only the way an unusable one + # is reported differs. A negative is rejected here, where the previous + # `max(0.0, …)` clamped it. + amount = coerce_rate(value) + return amount if amount is not None else 0.0 def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: @@ -326,9 +346,7 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: # actually deducts from the balance. For non-BYOK providers (e.g. # OpenRouter) usage.cost already equals upstream_inference_cost, so we # fall through to the normal ``cost`` lookup below. - upstream_cost = _coerce_usd( - cost_details.get("upstream_inference_cost") - ) + upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost")) if upstream_cost > 0 and usage_data.get("is_byok"): byok_fee = _coerce_usd(usage_data.get("cost")) return upstream_cost + byok_fee @@ -359,8 +377,7 @@ def _get_pricing_rates( ``None`` means configured fixed pricing should be used by the caller. """ if settings.fixed_pricing and ( - settings.fixed_per_1k_input_tokens - or settings.fixed_per_1k_output_tokens + settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens ): return None @@ -416,12 +433,8 @@ def _get_pricing_rates( usd_per_sat = sats_usd_price() mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat - cache_read_usd = _coerce_usd( - pricing.get("cache_read_input_token_cost") - ) - cache_write_usd = _coerce_usd( - pricing.get("cache_creation_input_token_cost") - ) + cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost")) + cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost")) mscr_1k = ( cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat if cache_read_usd > 0 @@ -479,6 +492,7 @@ def _calculate_from_usd_cost( """Calculate cost from USD figures, deriving input/output split from tokens.""" if provider_fee is None: provider_fee = _resolve_provider_fee(response_data.get("model", "")) + reported_usd = usd_cost usd_cost = usd_cost * provider_fee input_usd = input_usd * provider_fee output_usd = output_usd * provider_fee @@ -525,9 +539,7 @@ def _calculate_from_usd_cost( regular_weight = input_tokens * input_rate cache_read_weight = cache_read_tokens * cache_read_rate cache_creation_weight = cache_creation_tokens * cache_creation_rate - total_input_weight = ( - regular_weight + cache_read_weight + cache_creation_weight - ) + total_input_weight = regular_weight + cache_read_weight + cache_creation_weight if total_input_weight > 0: cache_read_msats = int( round( @@ -566,6 +578,7 @@ def _calculate_from_usd_cost( cache_creation_input_tokens=cache_creation_tokens, cache_read_msats=cache_read_msats, cache_creation_msats=cache_creation_msats, + upstream_usd=reported_usd, ) diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 702ab527..d088151f 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -1,8 +1,12 @@ +import asyncio import base64 +import ipaddress import json import math +import socket from io import BytesIO from typing import Any +from urllib.parse import urlsplit, urlunsplit import httpx from fastapi import HTTPException, Response @@ -14,10 +18,22 @@ from ..core import get_logger from ..core.exceptions import UpstreamError from ..core.redaction import redact_org_ids from ..core.settings import settings -from ..wallet import deserialize_token_from_string +from ..wallet import ( + UntrustedSourceMintError, + classify_redemption_error, + deserialize_token_from_string, + is_trusted_source_mint, +) +from .responses_input import ( + FILE_ID_URL_PREFIX, + count_input_images, + input_image_part_to_image_url, + responses_input_to_messages, +) logger = get_logger(__name__) + def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None: if x_cashu := headers.get("x-cashu", None): cashu_token = x_cashu @@ -68,6 +84,19 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N detail="Invalid authentication token format", ) + if not is_trusted_source_mint(token_obj.mint): + classified = classify_redemption_error( + UntrustedSourceMintError(f"Untrusted source mint: {token_obj.mint}") + ) + assert classified is not None + error_type, status_code, message, error_code = classified + raise HTTPException( + status_code=status_code, + detail={ + "error": {"message": message, "type": error_type, "code": error_code} + }, + ) + amount_msat = ( token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000 ) @@ -156,7 +185,12 @@ async def calculate_discounted_max_cost( body: dict, model_obj: Any | None = None, ) -> int: - """Calculate the discounted max cost for a request using model pricing when available.""" + """Calculate the discounted max cost for a request using model pricing when available. + + Completion discounts are trimmed from the largest declared cap among + ``max_tokens`` and ``max_completion_tokens`` (chat/completions) or + ``max_output_tokens`` (responses). + """ if settings.fixed_pricing: return max_cost_for_model @@ -196,41 +230,75 @@ async def calculate_discounted_max_cost( adjusted = max_cost_for_model - if messages := body.get("messages"): - prompt_tokens = estimate_tokens(messages) + messages = body.get("messages") + # Estimated over the whole body: a discount driven by message text alone lets + # a caller hide prompt weight elsewhere, shrink the reservation, and be billed + # for work the reservation never covered. + prompt_tokens = estimate_prompt_tokens(body) - image_tokens = await estimate_image_tokens_in_messages(messages) - if image_tokens > 0: - logger.debug( - "Found images in request", - extra={ - "model": model, - "image_tokens": image_tokens, - }, - ) - prompt_tokens += image_tokens + # Images are billed as tokens by the upstream but carry no text for + # ``estimate_prompt_tokens`` to count, so they are estimated separately and + # added on both the chat (``messages``) and Responses (``input``) paths. + image_tokens = 0 + if isinstance(messages, list): + image_tokens += await estimate_image_tokens_in_messages(messages) + input_data = body.get("input") + if isinstance(input_data, list): + converted = responses_input_to_messages(input_data) + if converted is None: + image_tokens += count_input_images(input_data) * _MAX_ORIGINAL_IMAGE_TOKENS + else: + image_tokens += await estimate_image_tokens_in_messages(converted) + if image_tokens > 0: + logger.debug( + "Found images in request", + extra={ + "model": model, + "image_tokens": image_tokens, + }, + ) + prompt_tokens += image_tokens + if prompt_tokens > 0: estimated_prompt_delta_sats = ( max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt ) if estimated_prompt_delta_sats > 0: adjusted = adjusted - math.floor(estimated_prompt_delta_sats * 1000) - max_tokens_raw = body.get("max_tokens", None) - if max_tokens_raw is not None: + # Completion caps arrive under several names: ``max_tokens`` (legacy + # chat), ``max_completion_tokens`` (modern chat) and ``max_output_tokens`` + # (Responses API). When a request declares more than one, reserve against + # the largest: upstream precedence between the fields varies by provider, + # so the smaller cap may not be honored and the reservation must never + # under-cover what the upstream could bill. + max_tokens_int: int | None = None + for cap_field in ("max_tokens", "max_completion_tokens", "max_output_tokens"): + cap_raw = body.get(cap_field) + if cap_raw is None: + continue try: - max_tokens_int = int(max_tokens_raw) + cap_int = int(cap_raw) except (TypeError, ValueError): logger.warning( - "Invalid max_tokens; ignoring in cost adjustment", - extra={"max_tokens": str(max_tokens_raw)[:64], "model": model}, + "Invalid completion token cap; ignoring in cost adjustment", + extra={ + "field": cap_field, + "value": str(cap_raw)[:64], + "model": model, + }, ) - else: - estimated_completion_delta_sats = ( - max_completion_allowed_sats - max_tokens_int * model_pricing.completion - ) - if estimated_completion_delta_sats > 0: - adjusted = adjusted - math.floor(estimated_completion_delta_sats * 1000) + continue + max_tokens_int = ( + cap_int if max_tokens_int is None else max(max_tokens_int, cap_int) + ) + + if max_tokens_int is not None: + estimated_completion_delta_sats = ( + max_completion_allowed_sats - max_tokens_int * model_pricing.completion + ) + if estimated_completion_delta_sats > 0: + adjusted = adjusted - math.floor(estimated_completion_delta_sats * 1000) logger.debug( "Discounted max cost computed", @@ -262,6 +330,51 @@ def estimate_tokens(messages: list) -> int: return total // 3 +def _sum_string_chars(node: Any) -> int: + """Recursively sum the length of every string in the tree, keys included. + + Nothing is excluded. Keys count because JSON-schema property names are + forwarded to the provider, and no exclusion rule can be trusted here: every + part of the body is caller-controlled, so any carve-out (by key name or by + value shape) is a place to hide prompt weight for free. Inline image data is + therefore counted as text too, which only makes the discount smaller. + """ + if isinstance(node, str): + return len(node) + if isinstance(node, dict): + return sum( + len(str(key)) + _sum_string_chars(value) for key, value in node.items() + ) + if isinstance(node, list): + return sum(_sum_string_chars(item) for item in node) + return 0 + + +def _count_prompt_token_ids(node: Any) -> int: + if isinstance(node, int) and not isinstance(node, bool): + return 1 + if isinstance(node, list): + return sum(_count_prompt_token_ids(item) for item in node) + return 0 + + +def estimate_prompt_tokens(body: dict) -> int: + """Conservatively estimate prompt tokens for the whole provider-bound body. + + Every string counts, as do token IDs in legacy ``prompt`` arrays, so no + forwarded field can hide prompt weight and shrink its reservation. + """ + return _sum_string_chars(body) // 3 + _count_prompt_token_ids(body.get("prompt")) + + +IMAGE_FETCH_TIMEOUT_SECONDS = 10.0 +# Dimensions live in the header, so a prefix suffices and an endless body cannot +# pin memory. +IMAGE_FETCH_MAX_BYTES = 512 * 1024 +# Fetches are sequential, so an unbounded URL list is a request-time amplifier. +IMAGE_FETCH_MAX_PER_REQUEST = 8 + + def _get_image_dimensions(image_data: bytes) -> tuple[int, int]: """Extract image dimensions from image bytes.""" try: @@ -275,13 +388,81 @@ def _get_image_dimensions(image_data: bytes) -> tuple[int, int]: return (512, 512) -async def _fetch_image_from_url(url: str) -> bytes | None: - """Fetch image from URL.""" +def _is_blocked_address(address: str) -> bool: + """Allow only globally reachable addresses (RFC 6890).""" try: - async with httpx.AsyncClient(timeout=10.0) as client: - response = await client.get(url) - response.raise_for_status() - return response.content + ip = ipaddress.ip_address(address) + except ValueError: + return True + if isinstance(ip, ipaddress.IPv6Address): + # An embedded v4 address would otherwise smuggle a rejected target past + # the v6 checks. + for embedded in (ip.ipv4_mapped, ip.sixtofour): + if embedded is not None: + return _is_blocked_address(str(embedded)) + return not ip.is_global or ip.is_multicast + + +async def _validated_fetch_target(url: str) -> tuple[str, str]: + """Return the URL to request and its ``Host`` header. + + Cost estimation runs on the unauthenticated request body, so a caller can + otherwise aim the node at internal hosts. HTTP is rewritten to the resolved + address so the name cannot rebind between check and connect; HTTPS keeps its + hostname because certificate validation already binds the connection. + """ + parts = urlsplit(url) + if parts.scheme not in ("http", "https"): + raise ValueError(f"unsupported scheme: {parts.scheme or 'none'}") + host = parts.hostname + if not host: + raise ValueError("missing host") + + default_port = 443 if parts.scheme == "https" else 80 + port = parts.port or default_port + host_header = f"[{host}]" if ":" in host else host + if parts.port is not None: + host_header = f"{host_header}:{parts.port}" + + infos = await asyncio.get_running_loop().getaddrinfo( + host, port, proto=socket.IPPROTO_TCP + ) + if not infos: + raise ValueError("host did not resolve") + for info in infos: + if _is_blocked_address(str(info[4][0])): + raise ValueError("host resolves to a blocked address") + + if parts.scheme == "https": + return url, host_header + + family, _, _, _, sockaddr = infos[0] + address = str(sockaddr[0]) + pinned = f"[{address}]" if family == socket.AF_INET6 else address + if parts.port is not None: + pinned = f"{pinned}:{parts.port}" + return urlunsplit((parts.scheme, pinned, parts.path, parts.query, "")), host_header + + +async def _fetch_image_from_url(url: str) -> bytes | None: + """Fetch the leading bytes of an image, enough to read its dimensions.""" + try: + target, host_header = await _validated_fetch_target(url) + async with httpx.AsyncClient( + timeout=IMAGE_FETCH_TIMEOUT_SECONDS, follow_redirects=False + ) as client: + async with client.stream( + "GET", target, headers={"Host": host_header} + ) as response: + response.raise_for_status() + chunks: list[bytes] = [] + downloaded = 0 + async for chunk in response.aiter_bytes(): + chunks.append(chunk) + downloaded += len(chunk) + if downloaded >= IMAGE_FETCH_MAX_BYTES: + break + return b"".join(chunks)[:IMAGE_FETCH_MAX_BYTES] except Exception as e: logger.warning( "Failed to fetch image from URL", @@ -290,15 +471,46 @@ async def _fetch_image_from_url(url: str) -> bytes | None: return None +# Patch-based image pricing (OpenAI ``detail: "original"``): the image is +# covered with 32x32px patches and billed as ceil(patches * multiplier) +# tokens, with no 512px-tile downscaling. The API rejects images above +# 30,000 patches, so at the 1.2x multiplier documented for the +# original-capable model families (gpt-5.4/5.5/5.6) the worst case a +# single image can bill is 36,000 tokens. +_IMAGE_PATCH_PX = 32 +_MAX_IMAGE_PATCHES = 30_000 +_MAX_ORIGINAL_IMAGE_TOKENS = (_MAX_IMAGE_PATCHES * 6 + 4) // 5 # 36,000 + + +def _calculate_original_image_tokens(width: int, height: int) -> int: + """Estimate tokens for an image billed at ``detail: "original"``. + + Patch-based models cover the image with 32x32px patches and bill + ``ceil(patches * 1.2)`` tokens. The estimate is bounded by the + 30,000-patch rejection limit, which is more conservative than the + per-model resizing patch budgets (e.g. 10,000 patches on gpt-5.4/5.5) + so it never under-reserves. + """ + patches = ((width + _IMAGE_PATCH_PX - 1) // _IMAGE_PATCH_PX) * ( + (height + _IMAGE_PATCH_PX - 1) // _IMAGE_PATCH_PX + ) + bounded = min(patches, _MAX_IMAGE_PATCHES) + return (bounded * 6 + 4) // 5 # ceil(bounded * 1.2) in exact integer math + + def _calculate_image_tokens(width: int, height: int, detail: str = "auto") -> int: """Calculate image tokens based on OpenAI's vision pricing. For low detail: 85 tokens For high detail/auto: 85 base tokens + 170 tokens per 512px tile + For original detail: patch-based pricing at the original resolution """ if detail == "low": return 85 + if detail == "original": + return _calculate_original_image_tokens(width, height) + if width > 2048 or height > 2048: aspect_ratio = width / height if width > height: @@ -330,6 +542,7 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: Supports both base64 encoded images and image URLs. """ total_image_tokens = 0 + fetches = 0 for message in messages: if not isinstance(message, dict): @@ -350,7 +563,9 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: continue content_type = content_item.get("type") - if content_type not in ("image_url", "input_image"): + if content_type == "input_image": + content_item = input_image_part_to_image_url(content_item) + elif content_type != "image_url": continue image_url_data = content_item.get("image_url") @@ -362,7 +577,7 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: detail = "auto" elif isinstance(image_url_data, dict): url = image_url_data.get("url", "") - detail = image_url_data.get("detail", "auto") + detail = image_url_data.get("detail") or "auto" else: continue @@ -370,49 +585,67 @@ async def estimate_image_tokens_in_messages(messages: list) -> int: continue if url.startswith("data:image/"): - try: - header, base64_data = url.split(",", 1) - image_bytes = base64.b64decode(base64_data) - width, height = _get_image_dimensions(image_bytes) - tokens = _calculate_image_tokens(width, height, detail) - total_image_tokens += tokens - logger.debug( - "Calculated tokens for base64 image", - extra={ - "width": width, - "height": height, - "detail": detail, - "tokens": tokens, - }, - ) - except Exception as e: - logger.warning( - "Failed to process base64 image", - extra={"error": str(e)}, - ) - total_image_tokens += 85 + total_image_tokens += _data_url_image_tokens(url, detail) + elif url.startswith(FILE_ID_URL_PREFIX): + total_image_tokens += _worst_case_image_tokens(detail) + elif fetches >= IMAGE_FETCH_MAX_PER_REQUEST: + logger.warning( + "Skipping image URL fetch above per-request limit", + extra={"url": url[:100], "limit": IMAGE_FETCH_MAX_PER_REQUEST}, + ) + total_image_tokens += _worst_case_image_tokens(detail) else: + fetches += 1 image_bytes_or_none = await _fetch_image_from_url(url) - if image_bytes_or_none: - width, height = _get_image_dimensions(image_bytes_or_none) - tokens = _calculate_image_tokens(width, height, detail) - total_image_tokens += tokens - logger.debug( - "Calculated tokens for URL image", - extra={ - "url": url[:100], - "width": width, - "height": height, - "detail": detail, - "tokens": tokens, - }, - ) - else: - total_image_tokens += 85 + total_image_tokens += _image_bytes_tokens( + image_bytes_or_none, detail, source=url[:100] + ) return total_image_tokens +def _worst_case_image_tokens(detail: str) -> int: + """Dimensions unknown: reserve the most ``detail`` can bill.""" + if detail == "original": + return _MAX_ORIGINAL_IMAGE_TOKENS + return _calculate_image_tokens(2048, 2048, detail) + + +def _data_url_image_tokens(url: str, detail: str) -> int: + try: + _, base64_data = url.split(",", 1) + image_bytes = base64.b64decode(base64_data, validate=True) + except Exception as e: + logger.warning("Failed to decode base64 image", extra={"error": str(e)}) + return _worst_case_image_tokens(detail) + return _image_bytes_tokens(image_bytes, detail, source="data-url") + + +def _image_bytes_tokens(image_bytes: bytes | None, detail: str, source: str) -> int: + if not image_bytes: + return _worst_case_image_tokens(detail) + try: + width, height = Image.open(BytesIO(image_bytes)).size + except Exception as e: + logger.warning( + "Failed to read image dimensions", + extra={"error": str(e), "source": source}, + ) + return _worst_case_image_tokens(detail) + tokens = _calculate_image_tokens(width, height, detail) + logger.debug( + "Calculated image tokens", + extra={ + "source": source, + "width": width, + "height": height, + "detail": detail, + "tokens": tokens, + }, + ) + return tokens + + def create_error_response( error_type: str, message: str, diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index 37ecb1cf..d5fb7e89 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -1,19 +1,28 @@ from __future__ import annotations -import math +import asyncio +import ipaddress +import json +import socket from collections.abc import Awaitable, Callable -from typing import TypedDict +from typing import Any, TypedDict import httpx from cashu.core.base import MeltQuoteState +from cashu.core.settings import settings as cashu_settings from cashu.wallet.wallet import Proof, Wallet +from ..cashu_compat import install_cashu_httpx_shim from ..mint import ( - MINT_TRANSPORT_EXCEPTIONS, is_mint_rate_limited, + is_mint_transport_error, run_mint_operation, ) +# cashu 0.20.x passes the `proxies` kwarg httpx removed in 0.28; see the module +# docstring. Installed at import so no mint call can run before the patch. +install_cashu_httpx_shim() + try: from bech32 import bech32_decode, convertbits # type: ignore except ModuleNotFoundError: # pragma: no cover – allow runtime miss @@ -42,6 +51,114 @@ class MeltOutcomeAmbiguousError(LNURLError): """ +class MeltUnpaidError(LNURLError): + """The mint answered the melt request itself with ``unpaid``. + + Unlike :class:`MeltOutcomeAmbiguousError` this is proof that no Lightning + payment was made, so callers may restore what they debited. + """ + + +_MAX_LNURL_REDIRECTS = 3 +_MAX_LNURL_RESPONSE_BYTES = 64 * 1024 +_NON_PUBLIC_HOST_SUFFIXES = (".localhost", ".local", ".internal") + + +async def _require_public_https_destination(url: httpx.URL) -> None: + """Reject anything that is not a public HTTPS endpoint. + + LNURL destinations and their redirect targets are attacker-influenced, so + every hop has to be re-checked: a single ``https://`` origin says nothing + about where a 302 points. A bare hostname check is not enough either: a + public-looking name can resolve to a loopback/link-local/private address + (SSRF), so DNS is resolved here and every resulting address must be global. + """ + if url.scheme != "https": + raise LNURLError("LNURL destination must be an HTTPS URL") + + host = (url.host or "").rstrip(".").lower() + if not host: + raise LNURLError("LNURL destination has no host") + + try: + literal = ipaddress.ip_address(host) + except ValueError: + literal = None + + if literal is not None: + if not literal.is_global: + raise LNURLError("LNURL destination is not a public host") + return + + if host == "localhost" or host.endswith(_NON_PUBLIC_HOST_SUFFIXES): + raise LNURLError("LNURL destination is not a public host") + + port = url.port or 443 + try: + infos = await asyncio.get_running_loop().getaddrinfo( + host, port, proto=socket.IPPROTO_TCP + ) + except socket.gaierror as e: + raise LNURLError("LNURL destination could not be resolved") from e + if not infos: + raise LNURLError("LNURL destination could not be resolved") + for info in infos: + try: + resolved = ipaddress.ip_address(info[4][0]) + except ValueError as e: + raise LNURLError("LNURL destination resolved to an invalid address") from e + if not resolved.is_global: + raise LNURLError("LNURL destination is not a public host") + + +async def _fetch_lnurl_json( + url: str, params: dict[str, int] | None = None +) -> dict[str, Any]: + """GET an LNURL endpoint, validating the destination at every redirect. + + Response bodies are never echoed: an LNURL service is untrusted, and its + payload would otherwise reach operator logs through raised errors. + """ + try: + target = httpx.URL(url, params=params) if params else httpx.URL(url) + except httpx.InvalidURL as e: + raise LNURLError("LNURL destination is not a usable URL") from e + await _require_public_https_destination(target) + + raw: bytes | None = None + async with httpx.AsyncClient() as client: + for _ in range(_MAX_LNURL_REDIRECTS + 1): + async with client.stream( + "GET", target, follow_redirects=False, timeout=10 + ) as response: + if response.is_redirect: + target = target.join(response.headers.get("location", "")) + await _require_public_https_destination(target) + continue + response.raise_for_status() + chunks = bytearray() + async for chunk in response.aiter_bytes(): + chunks.extend(chunk) + if len(chunks) > _MAX_LNURL_RESPONSE_BYTES: + raise LNURLError("LNURL response exceeded the size limit") + raw = bytes(chunks) + break + else: + raise LNURLError("LNURL destination exceeded the redirect limit") + + if raw is None: + raise LNURLError("LNURL destination exceeded the redirect limit") + + try: + data = json.loads(raw) + except ValueError as e: + raise LNURLError("LNURL response was not valid JSON") from e + + if not isinstance(data, dict): + raise LNURLError("LNURL response was not a JSON object") + return data + + async def decode_lnurl(lnurl: str) -> str: """Decode LNURL to get the actual URL. @@ -111,26 +228,30 @@ async def get_lnurl_data(lnurl: str) -> LNURLData: httpx.HTTPError: If the HTTP request fails """ url = await decode_lnurl(lnurl) - - async with httpx.AsyncClient() as client: - response = await client.get(url, follow_redirects=True, timeout=10) - response.raise_for_status() - - lnurl_data = response.json() + lnurl_data = await _fetch_lnurl_json(url) # Validate payRequest data if lnurl_data.get("tag") != "payRequest": - raise LNURLError( - f"Invalid LNURL tag: expected 'payRequest', got '{lnurl_data.get('tag')}'" - ) + raise LNURLError("Invalid LNURL tag: expected 'payRequest'") - if not isinstance(lnurl_data.get("callback"), str): + callback_url = lnurl_data.get("callback") + if not isinstance(callback_url, str): raise LNURLError("Invalid LNURL payRequest: missing callback URL") + try: + callback_target = httpx.URL(callback_url) + except httpx.InvalidURL as e: + raise LNURLError("Invalid LNURL callback URL") from e + await _require_public_https_destination(callback_target) + + min_sendable = lnurl_data.get("minSendable", 1000) # Default 1 sat + max_sendable = lnurl_data.get("maxSendable", 1000000000) # Default 1000 BTC + if not isinstance(min_sendable, int) or not isinstance(max_sendable, int): + raise LNURLError("Invalid LNURL payRequest: non-integer sendable limits") return LNURLData( - callback_url=lnurl_data["callback"], - min_sendable=lnurl_data.get("minSendable", 1000), # Default 1 sat - max_sendable=lnurl_data.get("maxSendable", 1000000000), # Default 1000 BTC + callback_url=callback_url, + min_sendable=min_sendable, + max_sendable=max_sendable, ) @@ -150,26 +271,53 @@ async def get_lnurl_invoice( LNURLError: If the response is invalid httpx.HTTPError: If the HTTP request fails """ - async with httpx.AsyncClient() as client: - response = await client.get( - callback_url, - params={"amount": amount_msat}, - follow_redirects=True, - timeout=10, - ) - response.raise_for_status() + invoice_data = await _fetch_lnurl_json(callback_url, params={"amount": amount_msat}) - invoice_data = response.json() - - if "pr" not in invoice_data: - # Check if there's an error in the response - if "reason" in invoice_data: - raise LNURLError(f"LNURL error: {invoice_data['reason']}") - raise LNURLError(f"Invalid LNURL invoice response: {invoice_data}") + if not isinstance(invoice_data.get("pr"), str): + raise LNURLError("LNURL callback returned no invoice") return invoice_data["pr"], invoice_data +def _select_melt_proofs( + wallet: Wallet, + proofs: list[Proof], + *, + quote_amount: int, + fee_reserve: int, + gross_budget: int, +) -> tuple[list[Proof] | None, int]: + """Select proofs that cover the quote and exact NUT-02 input fees. + + Cashu 0.20's ``select_to_send`` may recursively swap when asked to spend a + wallet's full balance. Melts accept overpayment and return change, so a + bounded, largest-first selection is both safer and minimizes input fees. + + Mints reject a melt carrying more than ``mint_max_request_length`` inputs, + so a dust-heavy wallet can only pay what its largest inputs cover; the + caller lowers the amount and the rest goes out on later payouts. + """ + selected: list[Proof] = [] + selected_amount = 0 + required = quote_amount + fee_reserve + spendable = [ + proof + for proof in sorted(proofs, key=lambda item: item.amount, reverse=True) + if getattr(proof, "reserved", False) is not True + ] + for proof in spendable[: cashu_settings.mint_max_request_length]: + selected.append(proof) + selected_amount += proof.amount + input_fees = int(wallet.get_fees_for_proofs(selected)) + required = quote_amount + fee_reserve + input_fees + if selected_amount >= required: + if required <= gross_budget: + return selected, 0 + # Covered but over budget; more proofs only raise input fees. + break + return None, max(1, required - min(selected_amount, gross_budget)) + + async def raw_send_to_lnurl( wallet: Wallet, proofs: list[Proof], @@ -201,12 +349,11 @@ async def raw_send_to_lnurl( # Send USD to Lightning Address paid = await wallet.send_to_lnurl("user@getalby.com", 50, unit="usd") """ - total_balance = sum(proof.amount for proof in proofs) - if amount and total_balance < amount: + if not isinstance(amount, int) or isinstance(amount, bool) or amount <= 0: + raise ValueError("A positive integer amount is required to send to an LNURL.") + if sum(proof.amount for proof in proofs) < amount: raise ValueError("Amount to send is higher than available proofs.") - else: - assert isinstance(amount, int) - total_balance = amount + total_balance = amount lnurl_data = await get_lnurl_data(lnurl) if unit == "sat": @@ -226,25 +373,51 @@ async def raw_send_to_lnurl( f"({min_sendable_sat} - {max_sendable_sat} {unit})" ) - estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + 1 - estimated_fees_msat = estimated_fees_sat * 1000 - final_amount = amount_msat - estimated_fees_msat + final_amount = amount_msat - bolt11_invoice, _ = await get_lnurl_invoice( - lnurl_data["callback_url"], final_amount - ) + selected_proofs: list[Proof] | None = None + # Find the largest amount covered by the budget after reserve and input fees. + for _ in range(8): + if final_amount < lnurl_data["min_sendable"]: + raise LNURLError("Cashu melt fees leave no payable LNURL amount") + bolt11_invoice, _ = await get_lnurl_invoice( + lnurl_data["callback_url"], final_amount + ) + melt_quote_resp = await run_mint_operation( + lambda: wallet.melt_quote(invoice=bolt11_invoice), + op_name="lnurl_melt_quote", + mint_url=str(wallet.url), + # Quote creation is unsafe to retry without idempotency. + retry_timeouts=False, + ) - melt_quote_resp = await run_mint_operation( - lambda: wallet.melt_quote(invoice=bolt11_invoice), - op_name="lnurl_melt_quote", - mint_url=str(wallet.url), - ) + quoted_amount = int(melt_quote_resp.amount) + expected_amount = final_amount // 1000 if unit == "sat" else final_amount + if quoted_amount != expected_amount: + raise LNURLError( + f"LNURL invoice amount does not match the requested amount " + f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})" + ) + + selected_proofs, shortfall = _select_melt_proofs( + wallet, + proofs, + quote_amount=quoted_amount, + fee_reserve=int(melt_quote_resp.fee_reserve), + gross_budget=amount, + ) + if selected_proofs is not None: + break + final_amount -= shortfall * (1000 if unit == "sat" else 1) + else: + raise LNURLError("Cashu melt fees exceed the requested gross amount") if on_melt_quote is not None: await on_melt_quote(melt_quote_resp.quote) - if amount: - proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) + assert selected_proofs is not None + proofs = selected_proofs + await wallet.set_reserved_for_send(proofs, reserved=True) try: melt_response = await run_mint_operation( @@ -265,15 +438,29 @@ async def raw_send_to_lnurl( # reserved as though a Lightning payment could still settle. await wallet.set_reserved_for_send(proofs, reserved=False) raise - if not isinstance(error, MINT_TRANSPORT_EXCEPTIONS): + if not is_mint_transport_error(error): raise + # Cashu clears reservations on transport errors despite an unknown outcome. + try: + await wallet.set_reserved_for_melt( + proofs, reserved=True, quote_id=melt_quote_resp.quote + ) + except Exception as reservation_error: + raise MeltOutcomeAmbiguousError( + "Melt outcome is ambiguous and its proof reservation could not " + "be restored; proofs must not be retried" + ) from reservation_error melt_response = None melt_error: BaseException | None = error else: melt_error = None - if getattr(melt_response, "state", None) == MeltQuoteState.paid: + melt_state = getattr(melt_response, "state", None) + if melt_state == MeltQuoteState.paid: return final_amount + if melt_state == MeltQuoteState.unpaid: + await wallet.set_reserved_for_send(proofs, reserved=False) + raise MeltUnpaidError("Cashu mint confirmed that the melt was unpaid") try: quote = await run_mint_operation( @@ -281,6 +468,8 @@ async def raw_send_to_lnurl( op_name="reconcile_lnurl_melt_quote", mint_url=str(wallet.url), retry_timeouts=False, + # Reconciliation must bypass the cooldown opened by this failure. + allow_during_cooldown=True, ) except Exception as reconciliation_error: raise MeltOutcomeAmbiguousError( @@ -290,6 +479,20 @@ async def raw_send_to_lnurl( if quote is not None and quote.state == MeltQuoteState.paid: return final_amount + if quote is not None and quote.state == MeltQuoteState.unpaid: + # A just-dispatched quote can briefly report unpaid before transitioning. + try: + await wallet.set_reserved_for_melt( + proofs, reserved=True, quote_id=melt_quote_resp.quote + ) + except Exception as reservation_error: + raise MeltOutcomeAmbiguousError( + "Melt outcome is ambiguous and its proof reservation could not " + "be restored; proofs must not be retried" + ) from reservation_error + raise MeltOutcomeAmbiguousError( + "Melt outcome is ambiguous; an immediate unpaid state is not final" + ) from melt_error state = getattr(getattr(quote, "state", None), "value", "unknown") raise MeltOutcomeAmbiguousError( diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 3deedb5a..bdd8fa73 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -5,13 +5,14 @@ import random import httpx from fastapi import APIRouter, Depends, HTTPException, Request from pydantic import BaseModel as V2BaseModel -from pydantic.v1 import BaseModel +from pydantic.v1 import BaseModel, validator from sqlmodel.ext.asyncio.session import AsyncSession from ..core.db import ModelRow, UpstreamProviderRow, get_session from ..core.logging import get_logger from ..core.settings import settings from .price import sats_usd_price +from .rates import BILLABLE_PRICING_FIELDS, coerce_rate, is_usable_rate logger = get_logger(__name__) @@ -58,12 +59,56 @@ class Pricing(BaseModel): max_cost: float = 0.0 # in sats not msats +# The rates ``Pricing`` declares without a default, derived from the model so the +# two cannot drift. A payload that omits one writes a row that will not parse. +REQUIRED_PRICING_FIELDS = tuple( + name for name, field in Pricing.__fields__.items() if field.required +) + + +def has_usable_pricing(pricing: Pricing) -> bool: + """True if every billable rate is a number a request could be billed on. + + Free is usable — a rate of zero is a real price. This asks only whether the + price is well-formed. One unusable rate disqualifies the whole price even + alongside a valid one, since a request can bill on the bad field: a positive + ``completion`` must not hide a negative ``prompt``. + """ + return all( + is_usable_rate(getattr(pricing, field)) for field in BILLABLE_PRICING_FIELDS + ) + + class TopProvider(BaseModel): context_length: int | None = None max_completion_tokens: int | None = None is_moderated: bool | None = None +class Reasoning(BaseModel): + """Per-model reasoning-effort metadata, matching OpenRouter's shape.""" + + mandatory: bool | None = None + default_enabled: bool | None = None + supported_efforts: list[str] | None = None + default_effort: str | None = None + supports_max_tokens: bool | None = None + + class Config: + extra = "ignore" + + def is_empty(self) -> bool: + return not any( + ( + self.mandatory is not None, + self.default_enabled is not None, + self.supported_efforts, + self.default_effort, + self.supports_max_tokens is not None, + ) + ) + + class Model(BaseModel): id: str name: str @@ -80,10 +125,43 @@ class Model(BaseModel): canonical_slug: str | None = None alias_ids: list[str] | None = None forwarded_model_id: str | None = None + reasoning: Reasoning | None = None + + class Config: + extra = "ignore" def __hash__(self) -> int: return hash(self.id) + @validator("reasoning", pre=True) + def _coerce_reasoning(cls, value: object) -> object: + if value is None or value is False: + return None + if isinstance(value, Reasoning): + return None if value.is_empty() else value + if not isinstance(value, dict) or not value: + return None + try: + parsed = Reasoning.parse_obj(value) + except Exception: + return None + return None if parsed.is_empty() else parsed + + def dict(self, **kwargs: object) -> dict: + # Non-reasoning models omit the field entirely so the catalog stays + # additive: existing clients never see a new null key. + data = super().dict(**kwargs) # type: ignore[arg-type] + reasoning = data.get("reasoning") + if not reasoning: + data.pop("reasoning", None) + elif isinstance(reasoning, dict): + cleaned = {k: v for k, v in reasoning.items() if v is not None} + if cleaned: + data["reasoning"] = cleaned + else: + data.pop("reasoning", None) + return data + def litellm_cost_entry(model_id: str) -> dict | None: """Look up ``model_id`` in litellm's bundled cost map. @@ -144,18 +222,17 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing: def _has_valid_pricing(model: dict) -> bool: - """Check if model has valid pricing (not free, no negative values).""" + """Check if model has valid pricing (usable rates, and not free).""" pricing = model.get("pricing", {}) if not pricing: return False - try: - prompt = float(pricing.get("prompt", 0)) - completion = float(pricing.get("completion", 0)) - except (ValueError, TypeError): - return False - - if prompt < 0 or completion < 0: + # Coercion runs before the both-zero test below, which `NaN` would defeat + # on its own — and one entry the coercion chokes on must not unwind the + # whole fetch, which once cost the node an entire upstream catalog. + prompt = coerce_rate(pricing.get("prompt", 0)) + completion = coerce_rate(pricing.get("completion", 0)) + if prompt is None or completion is None: return False if prompt == 0 and completion == 0: @@ -221,9 +298,10 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis return [] -def _row_to_model( +def _build_model_from_row( row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 ) -> Model: + """The deterministic USD view of a stored model row, before the sats conversion.""" architecture = json.loads(row.architecture) pricing = json.loads(row.pricing) per_request_limits = ( @@ -281,6 +359,14 @@ def _row_to_model( parsed_pricing.max_cost, ) = _calculate_usd_max_costs(model) + return model + + +def _row_to_model( + row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 +) -> Model: + model = _build_model_from_row(row, apply_provider_fee, provider_fee) + try: sats_to_usd = sats_usd_price() model = _update_model_sats_pricing(model, sats_to_usd) @@ -309,21 +395,57 @@ async def list_models( rows = (await session.exec(query)).all() # type: ignore provider_result = await session.exec(select(UpstreamProviderRow)) providers_by_id = {p.id: p for p in provider_result.all()} - return [ - _row_to_model( - r, - apply_provider_fee=apply_fees, - provider_fee=providers_by_id[r.upstream_provider_id].provider_fee - if r.upstream_provider_id in providers_by_id - else 1.01, - ) - for r in rows - if include_disabled - or ( + + models: list[Model] = [] + for r in rows: + if not include_disabled and not ( r.upstream_provider_id in providers_by_id and providers_by_id[r.upstream_provider_id].enabled - ) - ] + ): + continue + try: + model = _row_to_model( + r, + apply_provider_fee=apply_fees, + provider_fee=providers_by_id[r.upstream_provider_id].provider_fee + if r.upstream_provider_id in providers_by_id + else 1.01, + ) + except Exception as e: + # Stored pricing/architecture is JSON from whatever wrote the row, so + # a legacy import or foreign writer can leave a field that will not + # parse. Converting inside this loop meant one such row raised out of + # the whole listing and the node advertised nothing at all. Drop the + # row we cannot read — it is unservable either way — and keep serving + # the rest. + logger.warning( + "Skipping model row that could not be read", + extra={ + "model_id": r.id, + "upstream_provider_id": r.upstream_provider_id, + "error": str(e), + "error_type": type(e).__name__, + }, + ) + continue + # Served-map backstop for legacy rows and writers that bypass the admin + # edge: a negative or non-finite rate is not a price. Serving one + # advertises a rate the cost calculation cannot bill on, so the request + # falls through to the flat maximum reservation — or, if the rate is + # negative, bills an amount settlement credits back to the caller. + # ``include_disabled`` is the operator's listing, which must keep showing + # the row so it can be repaired. + if not include_disabled and not has_usable_pricing(model.pricing): + logger.warning( + "Withholding model with an unusable stored rate from the catalog", + extra={ + "model_id": r.id, + "upstream_provider_id": r.upstream_provider_id, + }, + ) + continue + models.append(model) + return models def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]: @@ -409,23 +531,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model: if (sats.max_cost or 0.0) < min_req_sats: sats.max_cost = min_req_sats - return Model( - id=model.id, - name=model.name, - created=model.created, - description=model.description, - context_length=model.context_length, - architecture=model.architecture, - pricing=model.pricing, - sats_pricing=sats, - per_request_limits=model.per_request_limits, - top_provider=model.top_provider, - enabled=model.enabled, - upstream_provider_id=model.upstream_provider_id, - canonical_slug=model.canonical_slug, - alias_ids=model.alias_ids, - forwarded_model_id=model.forwarded_model_id, - ) + return model.copy(update={"sats_pricing": sats}) except Exception as e: logger.error( "Failed to update sats pricing for model", diff --git a/routstr/payment/price.py b/routstr/payment/price.py index ad614322..e63d6754 100644 --- a/routstr/payment/price.py +++ b/routstr/payment/price.py @@ -5,12 +5,37 @@ import httpx from ..core import get_logger from ..core.settings import settings +from .rates import coerce_rate logger = get_logger(__name__) BTC_USD_PRICE: float | None = None SATS_USD_PRICE: float | None = None +SATS_PER_BTC = 100_000_000 + + +def _parse_quote(raw: object, exchange: str) -> float | None: + """Coerce an exchange quote to a price, or ``None`` if it is not one. + + Every quote passes through here because the aggregator takes the ``min()`` + of what it collects: an unusable quote does not merely join the sample, it + *wins* it, and the result is the rate every model and every request on the + node is priced at. A quote is stricter than a billable rate — it must be + positive, and positive *after* the sats conversion the node prices in: a + subnormal quote survives every guard here and still underflows to a zero + sats price, which then divides by zero on every model's rate. + """ + price = coerce_rate(raw) + if price is None or price <= 0 or price / SATS_PER_BTC <= 0: + logger.warning( + "Unusable price quote — ignoring this exchange", + extra={"exchange": exchange, "quote": repr(raw)}, + ) + return None + + return price + async def _kraken_btc_usd(client: httpx.AsyncClient) -> float | None: """Fetch BTC/USD price from Kraken API.""" @@ -18,10 +43,11 @@ async def _kraken_btc_usd(client: httpx.AsyncClient) -> float | None: try: response = await client.get(api) price_data = response.json() - price = float(price_data["result"]["XXBTZUSD"]["c"][0]) - - return price - except (httpx.RequestError, KeyError) as e: + return _parse_quote(price_data["result"]["XXBTZUSD"]["c"][0], "kraken") + except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e: + # A payload whose *shape* changed raises IndexError/TypeError, and a + # non-JSON body raises ValueError; unhandled, one exchange's bad day + # aborted the whole aggregation instead of dropping a single quote. logger.warning( "Kraken API error", extra={ @@ -39,10 +65,8 @@ async def _coinbase_btc_usd(client: httpx.AsyncClient) -> float | None: try: response = await client.get(api) price_data = response.json() - price = float(price_data["data"]["amount"]) - - return price - except (httpx.RequestError, KeyError) as e: + return _parse_quote(price_data["data"]["amount"], "coinbase") + except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e: logger.warning( "Coinbase API error", extra={ @@ -60,10 +84,8 @@ async def _binance_btc_usdt(client: httpx.AsyncClient) -> float | None: try: response = await client.get(api) price_data = response.json() - price = float(price_data["price"]) - - return price - except (httpx.RequestError, KeyError) as e: + return _parse_quote(price_data["price"], "binance") + except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e: logger.warning( "Binance API error", extra={ @@ -123,7 +145,7 @@ async def _update_prices() -> None: ) return BTC_USD_PRICE = btc_price - SATS_USD_PRICE = btc_price / 100_000_000 + SATS_USD_PRICE = btc_price / SATS_PER_BTC def btc_usd_price() -> float: diff --git a/routstr/payment/rates.py b/routstr/payment/rates.py new file mode 100644 index 00000000..998f3ca5 --- /dev/null +++ b/routstr/payment/rates.py @@ -0,0 +1,71 @@ +"""The one definition of a billable rate, with no dependencies of its own. + +A rate reaches the node from an upstream catalog, the LiteLLM cost map, an +operator's admin edit, a legacy database row and the BTC/USD feed. Each of those +readers needs the same two questions answered — is this value a rate at all, and +is it a rate a request can be billed on — so both answers live here, in a module +that imports nothing from the package. Every guard then shares one definition +instead of drifting, and no caller needs a deferred import to reach it. +""" + +import math + +# The rates a request can bill on. Derived fields (``max_*_cost``) are excluded +# — they are computed carriers, not charged rates. One definition, shared by the +# admin write edge and the served/routed guards, so they all cover the same set. +BILLABLE_PRICING_FIELDS = ( + "prompt", + "completion", + "request", + "image", + "web_search", + "internal_reasoning", + "input_cache_read", + "input_cache_write", +) + + +def is_usable_rate(rate: float) -> bool: + """True if a single billable rate is a number a request could be billed on. + + The one definition of a usable rate, so every guard that asks the question + answers it identically. A rate qualifies only when it is finite and + non-negative; zero is usable (it means "free", which is a real price) but + ``NaN``, ``±inf`` and negatives are not prices at all. + + Non-finite: ``inf > 0`` is True, so an infinite rate reads as chargeable and + would be served, routed and billed as ``inf``; ``NaN`` poisons every total it + enters and defeats ordinary comparisons, since ``NaN > 0``, ``NaN < 0`` and + ``NaN == 0`` are all False. Negative: a negative rate produces a negative + cost, which the settlement path subtracts from the balance — it pays the + caller to make requests. Both reach a stored row from upstream catalogs as + well as the admin edge (``json.loads`` accepts the bare ``NaN``/``Infinity`` + literals and overflows ``1e999`` to ``inf``). + + This is the rationale for every guard that calls it; the call sites say what + they do with the answer, not why the answer matters. + """ + return math.isfinite(rate) and rate >= 0.0 + + +def coerce_rate(value: object) -> float | None: + """Coerce a value from outside the node to a usable rate, or ``None``. + + The one coercion, shared by every reader of a rate the node did not compute + itself: an upstream catalog, the LiteLLM cost map, the exchange feed and the + admin write edge. Each of them was parsing for itself, and they disagreed — + which is how a boolean became a price on some paths and not others. + + A boolean is rejected outright: it is a change of shape, not a rate, and + Python would make ``True`` a finite, positive ``1.0`` that passes every + numeric guard downstream — a dollar per token. A numeric string is accepted, + because feeds report prices as strings. An oversized integer raises + ``OverflowError`` rather than ``ValueError``, so that is caught too. + """ + if isinstance(value, bool) or not isinstance(value, (int, float, str)): + return None + try: + rate = float(value) + except (TypeError, ValueError, OverflowError): + return None + return rate if is_usable_rate(rate) else None diff --git a/routstr/payment/responses_input.py b/routstr/payment/responses_input.py new file mode 100644 index 00000000..6a90a1e0 --- /dev/null +++ b/routstr/payment/responses_input.py @@ -0,0 +1,79 @@ +"""Convert a Responses API ``input`` into chat ``messages`` via litellm. + +litellm drops ``file_id`` (emits ``url: ""``) and nests a dict-form ``image_url`` +as-is, so ``input_image`` parts are flattened to ``{image_url: str, detail}`` first. +``file_id`` becomes a sentinel URL the image walker treats as unfetchable. +""" + +from typing import Any + +from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, +) + +from ..core import get_logger + +logger = get_logger(__name__) + +FILE_ID_URL_PREFIX = "file-id:" + + +def _flatten_input_image(part: dict[str, Any]) -> tuple[str, str]: + raw = part.get("image_url") + url = raw.get("url", "") if isinstance(raw, dict) else raw + detail = part.get("detail") or ( + raw.get("detail") if isinstance(raw, dict) else None + ) + if not url and part.get("file_id"): + url = f"{FILE_ID_URL_PREFIX}{part['file_id']}" + return (url if isinstance(url, str) else ""), (detail or "auto") + + +def _normalize_item(item: Any) -> Any: + if not isinstance(item, dict): + return item + if item.get("type") == "input_image": + url, detail = _flatten_input_image(item) + return {**item, "image_url": url, "detail": detail} + content = item.get("content") + if isinstance(content, list): + return {**item, "content": [_normalize_item(part) for part in content]} + return item + + +def input_image_part_to_image_url(part: dict[str, Any]) -> dict[str, Any]: + """Reshape an ``input_image`` part found inside chat ``messages``.""" + url, detail = _flatten_input_image(part) + return {"type": "image_url", "image_url": {"url": url, "detail": detail}} + + +def count_input_images(input_data: Any) -> int: + if isinstance(input_data, dict): + own = 1 if input_data.get("type") == "input_image" else 0 + return own + count_input_images(input_data.get("content")) + if isinstance(input_data, list): + return sum(count_input_images(item) for item in input_data) + return 0 + + +def responses_input_to_messages(input_data: Any) -> list[dict[str, Any]] | None: + """Returns ``None`` when the transform fails so the caller can worst-case.""" + if isinstance(input_data, str): + return [{"role": "user", "content": input_data}] + if not isinstance(input_data, list): + return [] + try: + normalized = [_normalize_item(item) for item in input_data] + converted = ( + LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( + input=normalized, # type: ignore[arg-type] + responses_api_request={}, + ) + ) + return [dict(message) for message in converted] + except Exception as e: + logger.warning( + "Responses input transform failed; using conservative image fallback", + extra={"error": str(e)}, + ) + return None diff --git a/routstr/proxy.py b/routstr/proxy.py index 4542d008..1f517779 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -26,6 +26,7 @@ from .core.db import ( ) from .core.exceptions import UpstreamError from .core.not_found import build_not_found_response +from .core.settings import settings from .payment.helpers import ( calculate_discounted_max_cost, check_token_balance, @@ -37,9 +38,18 @@ from .payment.models import Model from .upstream import BaseUpstreamProvider from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request from .upstream.helpers import init_upstreams +from .upstream.model_paths import ( + ModelPathSelector, + decode_model_path, + is_openrouter_base_url, + public_model_id, + public_provider_url, +) from .upstream.request_correction import correct_request, extract_error_message logger = get_logger(__name__) + +MODEL_PATH_HEADER = "x-routstr-model-path" proxy_router = APIRouter() _upstreams: list[BaseUpstreamProvider] = [] @@ -112,6 +122,25 @@ def get_candidates( return None +def _model_ids_match(requested: str, selected: str) -> bool: + if requested.lower() == selected.lower(): + return True + return public_model_id(requested).lower() == public_model_id(selected).lower() + + +def _candidate_for_selector( + selector: ModelPathSelector, + candidates: list[tuple[Model, BaseUpstreamProvider]], +) -> tuple[Model, BaseUpstreamProvider] | None: + for model_obj, upstream in candidates: + if ( + upstream.db_id == selector.provider_id + and public_provider_url(upstream.base_url) == selector.base_url + ): + return model_obj, upstream + return None + + def get_model_instance(model_id: str) -> Model | None: """Get the best-ranked Model instance for a model ID.""" candidates = get_candidates(model_id) @@ -221,22 +250,141 @@ async def refresh_model_maps_periodically() -> None: ) -_API_PATH_PREFIXES = ( - "v1/", - "responses", - "chat/", - "completions", - "models", - "embeddings", - "audio/", - "images/", - "moderations", - "providers", - "tee/", - "attestation", +# Canonical endpoints this proxy will forward, keyed by the path with any +# leading "v1/" and trailing slash removed, mapped to the methods allowed on +# each. The provider credential is attached during forwarding, so endpoint +# permission has to come from this table rather than from the client-supplied +# path: an upstream's key-management, organization, or billing routes live +# under the same origin and must never be reachable through the proxy. +_ALLOWED_ENDPOINTS: dict[str, frozenset[str]] = { + "chat/completions": frozenset({"POST"}), + "completions": frozenset({"POST"}), + "responses": frozenset({"POST"}), + "messages": frozenset({"POST"}), + "embeddings": frozenset({"POST"}), + "models": frozenset({"GET"}), + "attestation": frozenset({"GET"}), + "tee/attestation": frozenset({"GET"}), +} + +_ALLOWED_METHODS = frozenset({"GET", "POST"}) + + +def _canonical_api_path(path: str) -> str: + """Reduce a request path to its allowlist key. + + OpenAI-style clients reach the same endpoint with or without the ``v1/`` + prefix and with or without a trailing slash, so both spellings collapse to + one key. Callers must screen the path with + :func:`_is_ambiguously_spelled_path` first — this function assumes the path + has no dot segments, empty segments, or encoded separators left to resolve. + """ + core = path[:-1] if path.endswith("/") else path + if core.startswith("v1/"): + core = core[len("v1/") :] + return core + + +def _parse_extra_allowed_endpoints(raw: str) -> dict[str, frozenset[str]]: + """Parse operator-configured additions to the endpoint allowlist. + + Deployments whose provider exposes an endpoint outside the canonical set + opt in explicitly with ``PROXY_EXTRA_ALLOWED_PATHS``, a comma-separated + list of ``METHOD:path`` pairs (e.g. ``POST:v1/rerank,GET:batches``). Every + entry must name one concrete method and one unambiguous path; wildcards + and bare prefixes are deliberately unsupported, so widening the proxy's + reach is always a per-endpoint decision. Malformed entries are dropped + with a warning rather than silently widening or narrowing the surface. + """ + extra: dict[str, frozenset[str]] = {} + for entry in raw.split(","): + entry = entry.strip() + if not entry: + continue + method, separator, endpoint = entry.partition(":") + method = method.strip().upper() + endpoint = endpoint.strip() + if not separator or method not in _ALLOWED_METHODS or not endpoint: + logger.warning( + "Ignoring malformed PROXY_EXTRA_ALLOWED_PATHS entry", + extra={"entry": entry}, + ) + continue + if _is_ambiguously_spelled_path(endpoint): + logger.warning( + "Ignoring ambiguously spelled PROXY_EXTRA_ALLOWED_PATHS entry", + extra={"entry": entry}, + ) + continue + if any(character in endpoint for character in "*?["): + # Refuse glob syntax outright. Kept as a literal endpoint name it + # would never match a real request, so the operator would think + # they had widened the proxy when they had not. + logger.warning( + "Ignoring wildcard PROXY_EXTRA_ALLOWED_PATHS entry; " + "list each endpoint explicitly", + extra={"entry": entry}, + ) + continue + key = _canonical_api_path(endpoint) + extra[key] = extra.get(key, frozenset()) | {method} + return extra + + +def _is_ambiguously_spelled_path(path: str) -> bool: + """Reject paths whose spelling could resolve somewhere the allowlist did not. + + ``{path:path}`` arrives percent-decoded, so a client that sent ``%2e%2e`` or + ``%2f`` shows up here as ``..`` / ``/``. Dot segments, backslashes, duplicate + or leading separators, NUL bytes, and any residual encoded separator are + treated as unsafe: they let a caller walk off the canonical API surface (and + onto a sensitive upstream endpoint) even though the literal prefix check + would pass. Reject rather than trying to rewrite the path. + """ + if not path or path != path.strip() or path.startswith("/"): + return True + if "\x00" in path or "\\" in path: + return True + # A single trailing slash is canonical (e.g. "attestation/"); ignore it, + # then no remaining segment may be empty (covers "//") or a dot segment. + core = path[:-1] if path.endswith("/") else path + if any(segment in ("", ".", "..") for segment in core.split("/")): + return True + lowered = path.lower() + return "%2e" in lowered or "%2f" in lowered or "%5c" in lowered + + +_EXTRA_ALLOWED_ENDPOINTS = _parse_extra_allowed_endpoints( + settings.proxy_extra_allowed_paths ) +def _allowed_methods_for(endpoint: str) -> frozenset[str]: + """Return the methods allowed on a canonical endpoint, empty if unknown.""" + methods = _ALLOWED_ENDPOINTS.get(endpoint, frozenset()) + methods |= _EXTRA_ALLOWED_ENDPOINTS.get(endpoint, frozenset()) + return methods + + +def _forwarding_allowed(path: str, method: str) -> bool: + """Gate which method/path pairs may reach an upstream at all. + + The provider credential is attached during forwarding, so an unknown + endpoint must never be forwarded on the caller's say-so. The path is + reduced to its canonical form and looked up in the endpoint table; there is + no prefix match, so a known prefix no longer carries an unknown endpoint + (``v1/organization/api_keys`` is rejected even though ``v1/`` is familiar). + + EHBP requests are gated by the same table. Their body is opaque to the + proxy, which is a reason to constrain the destination more tightly, not to + trust the caller's path: the encrypted contract covers the body, never the + endpoint the credential is spent against. + """ + if method not in _ALLOWED_METHODS: + return False + return method in _allowed_methods_for(_canonical_api_path(path)) + + @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) @@ -255,14 +403,17 @@ async def proxy( async def _proxy( request: Request, path: str, session: AsyncSession ) -> Response | StreamingResponse: - # GET requests must hit a known API prefix; otherwise return a 404 (HTML - # for browsers, JSON for API clients). POST requests are always forwarded - # so that OpenAI-style endpoints work with or without the `v1/` prefix - # (e.g. `/chat/completions` as well as `/v1/chat/completions`). - if request.method == "GET" and not path.startswith(_API_PATH_PREFIXES): + # Screen the path before any routing decision: reject ambiguous spellings, + # then require a known API prefix so nothing unknown is forwarded with the + # provider credential attached. + if _is_ambiguously_spelled_path(path): return build_not_found_response(request, path) headers = dict(request.headers) + is_ehbp = "ehbp-encapsulated-key" in headers + + if not _forwarding_allowed(path, request.method): + return build_not_found_response(request, path) is_responses_api = path.startswith("v1/responses") or path.startswith("responses") request_body = await request.body() @@ -272,7 +423,6 @@ async def _proxy( # extract the model id, so the SDK sends it in X-Routstr-Model. Forward the # raw encrypted body to the upstream's /private/ endpoint and stream the # encrypted response back untouched — the SDK's SecureClient decrypts it. - is_ehbp = "ehbp-encapsulated-key" in headers if is_ehbp: request_body_dict = {} model_id = headers.get("x-routstr-model", "") @@ -294,6 +444,13 @@ async def _proxy( # without model/cost/auth lookups. Do not prefix-match here: paths such as # /attestationjunk must continue through normal authentication. if request.method == "GET" and _is_tinfoil_attestation_path(path): + if MODEL_PATH_HEADER in headers: + return create_error_response( + "unsupported_request", + "Model paths do not apply to attestation", + 400, + request=request, + ) selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams) if not selected_upstreams: return create_error_response( @@ -334,6 +491,41 @@ async def _proxy( "upstream_error", "All upstreams failed", 502, request=request ) + selector: ModelPathSelector | None = None + if MODEL_PATH_HEADER in headers: + selector = decode_model_path(headers[MODEL_PATH_HEADER]) + if ( + selector is None + or sum( + name.lower() == MODEL_PATH_HEADER for name, _ in request.headers.items() + ) + != 1 + ): + return create_error_response( + "invalid_request", + f"Malformed {MODEL_PATH_HEADER} header", + 400, + request=request, + ) + if not isinstance(model_id, str) or not _model_ids_match( + model_id, selector.model_id + ): + return create_error_response( + "invalid_request", + f"{MODEL_PATH_HEADER} selects model '{selector.model_id}' but the " + f"request asks for '{model_id}'", + 400, + request=request, + ) + if "models" in request_body_dict: + return create_error_response( + "invalid_request", + "Model paths cannot be combined with model fallbacks", + 400, + request=request, + ) + model_id = selector.model_id + candidates = get_candidates(model_id) if not candidates: @@ -341,6 +533,51 @@ async def _proxy( "invalid_model", f"Model '{model_id}' not found", 400, request=request ) + if selector is not None: + pinned = _candidate_for_selector(selector, candidates) + if pinned is None: + return create_error_response( + "invalid_model_path", + f"Model '{selector.model_id}' is not routable through provider " + f"{selector.provider_id}", + 404, + request=request, + ) + # Explicit routes must never enter cross-provider failover. + candidates = [pinned] + + if selector.endpoint_tag: + if ( + is_ehbp + or not request_body_dict + or not is_openrouter_base_url(pinned[1].base_url) + or _canonical_api_path(path) + not in {"chat/completions", "completions", "responses"} + ): + return create_error_response( + "unsupported_request", + "Endpoint pinning requires an OpenRouter completion or Responses JSON request", + 400, + request=request, + ) + provider_options = request_body_dict.get("provider", {}) + if not isinstance(provider_options, dict): + return create_error_response( + "invalid_request", + "provider must be an object", + 400, + request=request, + ) + request_body_dict = { + **request_body_dict, + "provider": { + **provider_options, + "order": [selector.endpoint_tag], + "allow_fallbacks": False, + }, + } + request_body = json.dumps(request_body_dict).encode() + if is_ehbp: candidates = [ (model, upstream) @@ -391,11 +628,21 @@ async def _proxy( ) elif is_responses_api: return await upstream.handle_x_cashu_responses( - request, x_cashu, path, max_cost_for_model, model_obj + request, + x_cashu, + path, + max_cost_for_model, + model_obj, + request_body=request_body, ) else: return await upstream.handle_x_cashu( - request, x_cashu, path, max_cost_for_model, model_obj + request, + x_cashu, + path, + max_cost_for_model, + model_obj, + request_body=request_body, ) except UpstreamError as e: logger.warning( @@ -593,17 +840,20 @@ async def _proxy( ) raise - # Reactive recovery: some models reject one specific request - # param (e.g. newer Anthropic models deprecating `temperature`). - # When the upstream 400s naming such a param, strip it from the - # body and retry the SAME upstream. ``already_stripped`` bounds - # this to one retry per distinct param so it always terminates. + # Same-provider recovery must not relax an explicit route. if response.status_code == 400 and not is_ehbp: correction = correct_request( request_body, extract_error_message(response), already_stripped, ) + if correction is not None and selector is not None: + corrected_body = json.loads(correction.body) + if any( + corrected_body.get(field) != request_body_dict.get(field) + for field in ("model", "provider") + ): + correction = None if correction is not None: request_body, bad_param = correction.body, correction.label already_stripped.add(bad_param) diff --git a/routstr/redemption_cache.py b/routstr/redemption_cache.py new file mode 100644 index 00000000..e7054b5b --- /dev/null +++ b/routstr/redemption_cache.py @@ -0,0 +1,109 @@ +"""In-memory negative cache for terminally failed Cashu token redemptions. + +A dead token (already spent, malformed, zero value) presented as a bearer key +triggers a full redemption attempt against the issuing mint on *every* request, +because no ``api_keys`` row survives the failed attempt. Polling clients that +never back off turn one dead token into thousands of pointless mint calls per +day. This cache remembers terminal redemption failures by token hash so +repeated presentations are rejected locally with the same error the mint +attempt would have produced. + +Only failures whose classification code is in :data:`TERMINAL_REDEMPTION_CODES` +are cached — transient failures (mint unreachable, rate-limited, cooldown) must +never be cached, or a brief mint outage would poison valid tokens. + +The cache is deliberately in-memory (bounded LRU with TTL) rather than a +database row: persisting a row per failed token would let an attacker fill the +database with garbage tokens for free. +""" + +import time +from collections import OrderedDict +from dataclasses import dataclass +from typing import Callable + +# Redemption ``code`` values that can never succeed on retry. A token that was +# already spent, failed to decode, or redeemed to zero value stays that way +# forever; swap fees exceeding the token amount only changes if the mint +# lowers its fees, which the TTL covers. +TERMINAL_REDEMPTION_CODES: frozenset[str] = frozenset( + { + "cashu_token_already_spent", + "invalid_cashu_token", + "cashu_token_zero_value", + "cashu_token_swap_fees_exceed_amount", + } +) + +DEFAULT_MAX_ENTRIES = 10_000 +DEFAULT_TTL_SECONDS = 24 * 60 * 60 + + +@dataclass(frozen=True) +class CachedRedemptionFailure: + """Sanitized classification of a terminal redemption failure. + + Mirrors the ``(type, status, message, code)`` tuple produced by + ``classify_redemption_error`` so a cache hit yields a byte-identical error + envelope to the original mint-backed failure. + """ + + status_code: int + error_type: str + message: str + code: str + + +class RedemptionNegativeCache: + """Bounded TTL+LRU cache keyed by the SHA-256 hash of the bearer token. + + Not thread-safe by design: all access happens on the asyncio event loop. + """ + + def __init__( + self, + max_entries: int = DEFAULT_MAX_ENTRIES, + ttl_seconds: float = DEFAULT_TTL_SECONDS, + clock: Callable[[], float] = time.monotonic, + ) -> None: + if max_entries <= 0: + raise ValueError("max_entries must be positive") + if ttl_seconds <= 0: + raise ValueError("ttl_seconds must be positive") + self._max_entries = max_entries + self._ttl_seconds = ttl_seconds + self._clock = clock + self._entries: OrderedDict[str, tuple[float, CachedRedemptionFailure]] = ( + OrderedDict() + ) + + def get(self, hashed_key: str) -> CachedRedemptionFailure | None: + entry = self._entries.get(hashed_key) + if entry is None: + return None + expires_at, failure = entry + if self._clock() >= expires_at: + del self._entries[hashed_key] + return None + self._entries.move_to_end(hashed_key) + return failure + + def put(self, hashed_key: str, failure: CachedRedemptionFailure) -> None: + if hashed_key in self._entries: + del self._entries[hashed_key] + elif len(self._entries) >= self._max_entries: + self._entries.popitem(last=False) + self._entries[hashed_key] = (self._clock() + self._ttl_seconds, failure) + + def discard(self, hashed_key: str) -> None: + self._entries.pop(hashed_key, None) + + def clear(self) -> None: + self._entries.clear() + + def __len__(self) -> int: + return len(self._entries) + + +# Process-wide singleton used by the bearer-auth path. +redemption_negative_cache = RedemptionNegativeCache() diff --git a/routstr/refund.py b/routstr/refund.py new file mode 100644 index 00000000..2aee20e8 --- /dev/null +++ b/routstr/refund.py @@ -0,0 +1,613 @@ +"""Refund claims: one open payout per API key, recorded before it is paid.""" + +import asyncio +import time +from typing import Any + +import httpx +from fastapi import HTTPException +from sqlalchemy.exc import IntegrityError +from sqlmodel import col, func, select, update + +from .core.db import ( + REFUND_OPEN_STATUSES, + REFUND_UNRESOLVED_STATUSES, + ApiKey, + AsyncSession, + Refund, + create_session, +) +from .core.db import ( + store_cashu_transaction_with_retry as store_cashu_transaction, +) +from .core.logging import get_logger +from .core.settings import settings +from .payment.lnurl import ( + LNURLError, + MeltOutcomeAmbiguousError, + MeltUnpaidError, + get_lnurl_data, +) +from .wallet import ( + check_bolt11_payment_status, + is_mint_connection_error, + send_to_lnurl, + send_token, + token_mint_url, +) + +logger = get_logger(__name__) + +RECONCILE_BATCH_LIMIT = 100 + + +def refund_unit(key: ApiKey) -> str: + return key.refund_currency or "sat" + + +def amount_in_unit(amount_msats: int, unit: str) -> int: + return amount_msats // 1000 if unit == "sat" else amount_msats + + +def refund_mint(key: ApiKey) -> str: + if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints: + return key.refund_mint_url + return settings.primary_mint + + +async def validate_lightning_destination(destination: str) -> None: + try: + await get_lnurl_data(destination) + except LNURLError as e: + raise HTTPException( + status_code=400, detail=f"Invalid lightning destination: {e}" + ) + except httpx.HTTPError as e: + raise HTTPException( + status_code=400, detail=f"Lightning destination unreachable: {e}" + ) + + +async def open_claim( + session: AsyncSession, + key: ApiKey, + *, + method: str, + destination: str | None, +) -> Refund: + """Zero the balance and insert the claim in one transaction.""" + unit = refund_unit(key) + # created_at has second resolution; step past the previous claim so the + # newest claim for a key always sorts first. + latest = await session.exec( + select(func.max(col(Refund.created_at))).where( + Refund.api_key_hashed_key == key.hashed_key + ) + ) + created_at = max(int(time.time()), (latest.one() or 0) + 1) + refund = Refund( + api_key_hashed_key=key.hashed_key, + method=method, + destination=destination, + amount_msats=key.total_balance, + unit=unit, + mint_url=refund_mint(key), + claimed_at=int(time.time()), + created_at=created_at, + updated_at=created_at, + ) + debit = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.balance) == key.balance) + .where(col(ApiKey.reserved_balance) == key.reserved_balance) + .values(balance=0, reserved_balance=0, reserved_at=None) + ) + try: + debited = await session.exec(debit) # type: ignore[call-overload] + if debited.rowcount == 0: + await session.rollback() + raise HTTPException( + status_code=409, + detail="Balance changed concurrently. Please retry the refund.", + ) + session.add(refund) + await session.commit() + except IntegrityError: + await session.rollback() + raise refund_in_progress_error() + return refund + + +async def _close( + session: AsyncSession, + refund: Refund, + *, + require_no_quote: bool = False, + require_no_token: bool = False, + from_statuses: tuple[str, ...] = REFUND_OPEN_STATUSES, + **values: object, +) -> bool: + stmt = ( + update(Refund) + .where(col(Refund.id) == refund.id) + .where(col(Refund.status).in_(from_statuses)) + ) + if require_no_quote: + # A quote recorded since the row was read means a melt may be in flight. + stmt = stmt.where(col(Refund.quote_id).is_(None)) + if require_no_token: + stmt = stmt.where(col(Refund.token).is_(None)) + result = await session.exec( # type: ignore[call-overload] + stmt.values(claimed_at=None, updated_at=int(time.time()), **values) + ) + return bool(result.rowcount) + + +async def renew_lease(session: AsyncSession, refund: Refund) -> None: + """Push the reconciler lease forward before a slow mint step.""" + await session.exec( # type: ignore[call-overload] + update(Refund) + .where(col(Refund.id) == refund.id) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + .values(claimed_at=int(time.time())) + ) + await session.commit() + + +async def record_quote(refund: Refund, quote_id: str, mint_url: str) -> None: + """Store the quote and its mint before the melt is sent; raises if the claim closed. + + Mint fallback can issue the quote on a different mint than the claim's. + """ + async with create_session() as session: + result = await session.exec( # type: ignore[call-overload] + update(Refund) + .where(col(Refund.id) == refund.id) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + # Renew the lease so the reconciler leaves the payout alone. + .values( + quote_id=quote_id, + mint_url=mint_url, + claimed_at=int(time.time()), + updated_at=int(time.time()), + ) + ) + await session.commit() + if not result.rowcount: + raise LNURLError("Refund claim closed before the melt was dispatched") + refund.quote_id = quote_id + refund.mint_url = mint_url + + +async def settle( + session: AsyncSession, + refund: Refund, + *, + quote_id: str | None = None, + token: str | None = None, + mint_url: str | None = None, +) -> bool: + values: dict[str, Any] = {"status": "paid"} + if quote_id is not None: + values["quote_id"] = quote_id + if token is not None: + values["token"] = token + if mint_url is not None: + values["mint_url"] = mint_url + # The payout side knows the money moved, so a claim the reconciler gave up + # on (stuck) is closed as paid too. + settled = await _close( + session, refund, from_statuses=REFUND_UNRESOLVED_STATUSES, **values + ) + await session.commit() + if not settled: + logger.warning( + "refund paid but its claim was already closed", + extra={"refund_id": refund.id, "prior_status": refund.status}, + ) + return settled + + +async def release( + session: AsyncSession, refund: Refund, *, require_no_quote: bool = False +) -> bool: + """Mark the claim failed and restore the balance.""" + if not await _close( + session, refund, require_no_quote=require_no_quote, status="failed" + ): + # Commit, not rollback: rollback after an ORM UPDATE breaks later loads. + await session.commit() + return False + await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == refund.api_key_hashed_key) + .values(balance=col(ApiKey.balance) + refund.amount_msats) + ) + await session.commit() + logger.info( + "refund released; balance restored", + extra={ + "refund_id": refund.id, + "key_hash": refund.api_key_hashed_key[:8], + "restored_msats": refund.amount_msats, + }, + ) + return True + + +async def hold( + session: AsyncSession, + refund: Refund, + quote_id: str | None, + *, + token: str | None = None, +) -> None: + """Withhold the balance; the quote or token names what the mint may have paid.""" + values: dict[str, Any] = {"status": "ambiguous", "quote_id": quote_id} + if token is not None: + values["token"] = token + values["mint_url"] = refund.mint_url + await _close(session, refund, **values) + await session.commit() + + +def refund_in_progress_error(refund: Refund | None = None) -> HTTPException: + """The 409 raised when a key already has an unresolved refund claim.""" + stuck = refund is not None and refund.status == "stuck" + error: dict[str, str] = { + "message": ( + "A refund for this key is unresolved and requires operator reconciliation." + if stuck + else "A refund for this key is already in progress." + ), + "type": "invalid_request_error", + "code": "refund_unresolved" if stuck else "refund_in_progress", + } + if refund is not None: + error["refund_id"] = refund.id + error["status"] = refund.status + return HTTPException(status_code=409, detail={"error": error}) + + +async def _latest_with_status( + session: AsyncSession, key: ApiKey, statuses: tuple[str, ...] +) -> Refund | None: + result = await session.exec( + select(Refund) + .where(Refund.api_key_hashed_key == key.hashed_key) + .where(col(Refund.status).in_(statuses)) + .order_by(col(Refund.created_at).desc(), col(Refund.updated_at).desc()) + ) + return result.first() + + +async def latest_open(session: AsyncSession, key: ApiKey) -> Refund | None: + """Latest non-terminal (in-flight) claim for the key, if any. + + An open claim means a prior refund already debited the balance and is still + settling, so the balance reads as zero even though a refund is under way. + """ + return await _latest_with_status(session, key, REFUND_OPEN_STATUSES) + + +async def latest_stuck(session: AsyncSession, key: ApiKey) -> Refund | None: + """Latest claim the reconciler gave up on; needs operator recovery.""" + return await _latest_with_status(session, key, ("stuck",)) + + +async def latest_terminal(session: AsyncSession, key: ApiKey) -> Refund | None: + """Latest paid refund of either method. + + Cashu tokens are normally served from cashu_transactions, which tracks + collection and sweeping; the claim row is the fallback when that ledger + write failed after the token was already issued. + """ + return await _latest_with_status(session, key, ("paid",)) + + +def describe(refund: Refund) -> dict[str, str]: + body: dict[str, str] = {"refund_id": refund.id, "status": refund.status} + if refund.token: + body["token"] = refund.token + if refund.destination: + body["recipient"] = refund.destination + if refund.unit == "sat": + body["sats"] = str(refund.amount_msats // 1000) + else: + body["msats"] = str(refund.amount_msats) + return body + + +async def _pay_lightning(session: AsyncSession, refund: Refund) -> bool: + async def capture_quote(quote: str, mint_url: str) -> None: + await record_quote(refund, quote, mint_url) + + try: + await send_to_lnurl( + amount_in_unit(refund.amount_msats, refund.unit), + refund.unit, + refund.mint_url, + str(refund.destination), + on_melt_quote=capture_quote, + ) + except MeltOutcomeAmbiguousError as e: + await hold(session, refund, refund.quote_id) + logger.error( + "refund outcome ambiguous; balance withheld pending reconciliation", + extra={ + "refund_id": refund.id, + "error": str(e), + "key_hash": refund.api_key_hashed_key[:8], + "quote_id": refund.quote_id, + }, + ) + raise + return await settle(session, refund, quote_id=refund.quote_id) + + +async def _pay_cashu(session: AsyncSession, refund: Refund) -> bool: + amount = amount_in_unit(refund.amount_msats, refund.unit) + await renew_lease(session, refund) + token = await send_token(amount, refund.unit, refund.mint_url) + # From here the token is bearer money: keep it on the claim so a failed + # settle withholds the balance instead of restoring it. + refund.token = token + refund.mint_url = token_mint_url(token, refund.mint_url) + return await settle(session, refund, token=token, mint_url=refund.mint_url) + + +async def _record_cashu_payout(refund: Refund) -> None: + """Ledger write for an issued token; the claim row already holds the token, + so a failure here must not fail the request or release the balance.""" + try: + await store_cashu_transaction( + token=str(refund.token), + amount=amount_in_unit(refund.amount_msats, refund.unit), + unit=refund.unit, + mint_url=refund.mint_url, + typ="out", + collected=False, + source="apikey", + api_key_hashed_key=refund.api_key_hashed_key, + ) + except Exception as e: + logger.error( + "refund token issued but cashu transaction was not recorded", + extra={ + "refund_id": refund.id, + "error": str(e), + "error_type": type(e).__name__, + "key_hash": refund.api_key_hashed_key[:8], + }, + ) + + +def unresolved_refund_error() -> HTTPException: + return HTTPException( + status_code=502, + detail=( + "Refund was dispatched but its outcome is unconfirmed; the " + "balance is withheld until reconciliation completes" + ), + ) + + +async def _abort(session: AsyncSession, refund: Refund) -> None: + """Fail the claim, or withhold it once a melt quote or token exists. + + A recorded quote means the mint may already have paid; an issued token is + already bearer money. In both cases the balance must not be restored. + """ + if refund.quote_id is None and refund.token is None: + await release(session, refund) + return + await hold(session, refund, refund.quote_id, token=refund.token) + logger.error( + "refund failed after its payout was dispatched; balance withheld " + "pending reconciliation", + extra={ + "refund_id": refund.id, + "key_hash": refund.api_key_hashed_key[:8], + "quote_id": refund.quote_id, + "has_token": refund.token is not None, + "mint_url": refund.mint_url, + }, + ) + raise unresolved_refund_error() + + +async def execute(session: AsyncSession, refund: Refund) -> dict[str, str]: + attached_refund = refund + # Keep payout evidence outside the identity map: a failed flush/commit can + # expire attached attributes, including the only copy of an issued token. + refund = Refund(**refund.model_dump()) + try: + if refund.method == "lightning": + settled = await _pay_lightning(session, refund) + else: + settled = await _pay_cashu(session, refund) + except MeltOutcomeAmbiguousError: + # Already held by _pay_lightning; releasing here would pay out twice. + raise unresolved_refund_error() + except MeltUnpaidError as e: + # The mint answered the melt itself with unpaid: proof that nothing was + # sent, so the balance goes back now rather than after reconciliation. + await release(session, refund) + logger.warning( + "refund melt unpaid at the mint; balance restored", + extra={ + "refund_id": refund.id, + "error": str(e), + "key_hash": refund.api_key_hashed_key[:8], + "quote_id": refund.quote_id, + }, + ) + raise HTTPException( + status_code=503, + detail="Lightning payment failed at the mint; balance restored. Retry later.", + ) + except HTTPException: + await session.rollback() + await _abort(session, refund) + raise + except Exception as e: + await session.rollback() + await _abort(session, refund) + logger.error( + "refund payout failed", + extra={ + "refund_id": refund.id, + "error": str(e), + "error_type": type(e).__name__, + "key_hash": refund.api_key_hashed_key[:8], + "method": refund.method, + "mint_url": refund.mint_url, + }, + ) + if is_mint_connection_error(e): + raise HTTPException(status_code=503, detail="Mint service unavailable") + raise HTTPException(status_code=500, detail="Refund failed") + + if refund.method == "cashu": + await _record_cashu_payout(refund) + + if settled: + refund.status = "paid" + refund.claimed_at = None + else: + # Report the row as it stands rather than a status that was not written. + await session.refresh(attached_refund) + refund = attached_refund + logger.info( + "refund paid", + extra={ + "refund_id": refund.id, + "method": refund.method, + "amount_msats": refund.amount_msats, + "key_hash": refund.api_key_hashed_key[:8], + }, + ) + return describe(refund) + + +async def _lease(refund_id: str, now: int, lease_cutoff: int) -> bool: + async with create_session() as session: + result = await session.exec( # type: ignore[call-overload] + update(Refund) + .where(col(Refund.id) == refund_id) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + .where( + col(Refund.claimed_at).is_(None) + | (col(Refund.claimed_at) < lease_cutoff) + ) + .values(claimed_at=now) + ) + await session.commit() + return bool(result.rowcount) + + +async def _reconcile(refund: Refund, now: int) -> None: + if refund.method != "lightning": + if refund.token is not None: + # The token was issued and kept on the claim; the payout is done. + async with create_session() as session: + await settle(session, refund) + return + # No quote to query for cashu; withhold the balance and alert once. + async with create_session() as session: + if await _close(session, refund, require_no_token=True, status="stuck"): + await session.commit() + logger.critical( + "cashu refund stuck; balance withheld, manual reconciliation required", + extra={ + "refund_id": refund.id, + "key_hash": refund.api_key_hashed_key[:8], + "amount_msats": refund.amount_msats, + }, + ) + else: + await session.commit() + current = await session.get(Refund, refund.id) + if current is not None and current.token is not None: + await settle(session, current) + return + + if refund.quote_id is None: + # Never sent, unless a quote appeared since the row was read. + async with create_session() as session: + if not await release(session, refund, require_no_quote=True): + logger.info( + "refund gained a melt quote during reconciliation; left open", + extra={"refund_id": refund.id}, + ) + return + + status = await check_bolt11_payment_status( + refund.mint_url, refund.unit, refund.quote_id + ) + if status == "paid": + async with create_session() as session: + await settle(session, refund) + elif status == "unpaid": + # A fresh melt can report unpaid briefly; trust it only after a timeout. + if refund.updated_at > now - settings.refund_claim_timeout_seconds: + logger.info( + "refund unpaid at the mint but too recent to release; waiting", + extra={"refund_id": refund.id, "updated_at": refund.updated_at}, + ) + return + async with create_session() as session: + await release(session, refund) + else: + logger.warning( + "refund still unresolved at the mint", + extra={"refund_id": refund.id, "melt_status": status}, + ) + + +async def reconcile_once() -> None: + """Resolve open claims whose lease has lapsed.""" + now = int(time.time()) + lease_cutoff = now - settings.refund_claim_timeout_seconds + async with create_session() as session: + result = await session.exec( + select(Refund) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + .where( + col(Refund.claimed_at).is_(None) + | (col(Refund.claimed_at) < lease_cutoff) + ) + .order_by(col(Refund.created_at)) + .limit(RECONCILE_BATCH_LIMIT) + ) + stale = list(result.all()) + + for refund in stale: + if not await _lease(refund.id, now, lease_cutoff): + continue + try: + await _reconcile(refund, now) + except Exception as e: + logger.error( + "refund reconciliation failed", + extra={ + "refund_id": refund.id, + "error": str(e), + "error_type": type(e).__name__, + }, + ) + + +async def periodic_refund_reconcile() -> None: + while True: + await asyncio.sleep(settings.refund_reconcile_interval_seconds) + try: + await reconcile_once() + except asyncio.CancelledError: + raise + except Exception as e: + logger.error( + "refund reconcile loop error", + extra={"error": str(e), "error_type": type(e).__name__}, + ) diff --git a/routstr/upstream/auto_topup.py b/routstr/upstream/auto_topup.py index c8eadb6e..a73bdecf 100644 --- a/routstr/upstream/auto_topup.py +++ b/routstr/upstream/auto_topup.py @@ -15,9 +15,6 @@ from ..core.db import ( UpstreamProviderRow, create_session, ) -from ..core.db import ( - store_cashu_transaction_with_retry as store_cashu_transaction, -) from ..payment.price import sats_usd_price from ..wallet import ( Bolt11PaymentAmbiguous, @@ -27,7 +24,7 @@ from ..wallet import ( maximum_owner_cashu_balance_sats, prepare_bolt11_payment, release_token_reservation, - send_token, + send_token_from_owner_locked, token_mint_url, wallet_operation_guard, ) @@ -49,6 +46,7 @@ PPQ_PHASES = frozenset({PPQ_PHASE_CLAIMED, PPQ_PHASE_IN_FLIGHT, PPQ_PHASE_RECONC PPQ_SETTLEMENT_ATTEMPTS = 5 PPQ_SETTLEMENT_POLL_SECONDS = 2 PPQ_PENDING_TTL_SECONDS = 15 * 60 +PPQ_SETTLED_COOLDOWN_SECONDS = 5 * 60 PPQ_MAX_INVOICE_PREMIUM = 1.10 PPQ_MIN_TOPUP_USD = 1 PPQ_MAX_TOPUP_USD = 500 @@ -58,6 +56,33 @@ PPQ_MAX_TOPUP_USD = 500 # the damage instead of letting the worker drain the owner's mint funds one # per-transaction-capped payment at a time. PPQ_MAX_DAILY_TOPUP_USD = 300 +# Routstr-to-Routstr claim lifecycle. "claimed" holds the slot while the token +# is being minted; nothing has left the wallet yet. "sent" means a bearer token +# was handed to the peer and only the peer's balance can say whether it landed. +# "backoff" holds the failure count between attempts, and "halted" stops the +# provider entirely until an admin releases it. +ROUTSTR_PHASE_CLAIMED = "claimed" +ROUTSTR_PHASE_SENT = "sent" +ROUTSTR_PHASE_BACKOFF = "backoff" +ROUTSTR_PHASE_HALTED = "halted" +ROUTSTR_PHASES = frozenset( + { + ROUTSTR_PHASE_CLAIMED, + ROUTSTR_PHASE_SENT, + ROUTSTR_PHASE_BACKOFF, + ROUTSTR_PHASE_HALTED, + } +) +ROUTSTR_PENDING_TTL_SECONDS = 15 * 60 +ROUTSTR_BACKOFF_BASE_SECONDS = 15 * 60 +ROUTSTR_MAX_TOPUP_FAILURES = 3 +ROUTSTR_MIN_TOPUP_SATS = 1 +ROUTSTR_MAX_TOPUP_SATS = 1_000_000 +# Rolling 24h ceiling on total Routstr auto top-up spend across all peers. The +# per-attempt claim already stops a peer from being paid twice for the same +# uncredited token; this bounds the total even when every attempt is credited +# and the peer simply keeps reporting a below-threshold balance. +ROUTSTR_MAX_DAILY_TOPUP_SATS = 2_000_000 async def periodic_auto_topup() -> None: @@ -72,7 +97,7 @@ async def periodic_auto_topup() -> None: except Exception as e: logger.error( "Auto top-up cycle failed", - extra={"error": str(e), "error_type": type(e).__name__}, + extra={"error": repr(e), "error_type": type(e).__name__}, ) await asyncio.sleep(AUTO_TOPUP_INTERVAL_SECONDS) @@ -103,7 +128,7 @@ async def _run_auto_topup_cycle() -> None: extra={ "provider_id": row.id, "base_url": row.base_url, - "error": str(e), + "error": repr(e), "error_type": type(e).__name__, }, ) @@ -134,12 +159,16 @@ async def _reconcile_all_ppq_claims() -> set[int]: except Exception as e: logger.error( "PPQ claim reconciliation failed", - extra={"provider_id": row.id, "error": str(e)}, + extra={ + "provider_id": row.id, + "error": repr(e), + "error_type": type(e).__name__, + }, ) return active_provider_ids -def _invalid_ppq_number(value: object, *, integer: bool = False) -> bool: +def _invalid_topup_number(value: object, *, integer: bool = False) -> bool: if isinstance(value, bool) or not isinstance(value, (int, float)): return True try: @@ -160,9 +189,9 @@ def validate_ppq_auto_topup_settings(settings: dict | None) -> str | None: threshold = settings.get("topup_threshold") amount = settings.get("topup_amount_limit") - if _invalid_ppq_number(threshold): + if _invalid_topup_number(threshold): return "PPQ auto top-up threshold must be a positive number" - if _invalid_ppq_number(amount, integer=True): + if _invalid_topup_number(amount, integer=True): return "PPQ auto top-up amount must be a positive whole number" amount_usd = int(typing.cast(int | float, amount)) if not PPQ_MIN_TOPUP_USD <= amount_usd <= PPQ_MAX_TOPUP_USD: @@ -173,6 +202,62 @@ def validate_ppq_auto_topup_settings(settings: dict | None) -> str | None: return None +_legacy_threshold_hinted: set[int] = set() + + +def _routstr_threshold_sats(row: UpstreamProviderRow, settings: dict) -> float: + """Balance, in sats, below which a top-up fires. + + ``topup_threshold_sats`` is used as written. A legacy ``topup_threshold`` + keeps the thousandfold it has always been compared with — reinterpreting it + as sats would drop an operator's trigger point by a factor of 1000 on + upgrade and leave the peer to run dry. The hint names the value to migrate + to, once per provider rather than once per scheduler tick. + """ + explicit = settings.get("topup_threshold_sats") + if explicit is not None: + return float(typing.cast(int | float, explicit)) + + threshold_sats = float(typing.cast(int | float, settings["topup_threshold"])) * 1000 + if row.id is not None and row.id not in _legacy_threshold_hinted: + _legacy_threshold_hinted.add(row.id) + logger.warning( + "Routstr auto top-up uses the legacy unitless threshold; set " + "topup_threshold_sats to state the unit", + extra={"provider_id": row.id, "threshold_sats": threshold_sats}, + ) + return threshold_sats + + +def validate_routstr_auto_topup_settings(settings: dict | None) -> str | None: + """Return why enabled Routstr auto top-up settings are invalid, if anything.""" + if not settings or not settings.get("auto_topup"): + return None + + threshold_sats = settings.get("topup_threshold_sats") + legacy_threshold = settings.get("topup_threshold") + amount = settings.get("topup_amount_limit") + mint_url = settings.get("topup_mint_url") + if threshold_sats is not None: + if _invalid_topup_number(threshold_sats): + return "Routstr auto top-up threshold must be a positive number of sats" + elif legacy_threshold is None: + return "Routstr auto top-up requires a threshold" + elif _invalid_topup_number(legacy_threshold): + return "Routstr auto top-up threshold must be a positive number" + if _invalid_topup_number(amount, integer=True): + return "Routstr auto top-up amount must be a positive whole number" + amount_sats = int(typing.cast(int | float, amount)) + if not ROUTSTR_MIN_TOPUP_SATS <= amount_sats <= ROUTSTR_MAX_TOPUP_SATS: + return ( + f"Routstr auto top-up amount must be between {ROUTSTR_MIN_TOPUP_SATS} " + f"and {ROUTSTR_MAX_TOPUP_SATS} sats" + ) + if not isinstance(mint_url, str) or not mint_url.strip(): + return "Routstr auto top-up requires a mint URL" + return None + + async def _check_and_topup_ppq_from_row(row: UpstreamProviderRow) -> None: settings: dict = {} if row.provider_settings: @@ -212,22 +297,18 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: if not settings.get("auto_topup"): return - threshold = settings.get("topup_threshold") - amount = settings.get("topup_amount_limit") - mint_url = settings.get("topup_mint_url") - - if not threshold or not amount or not mint_url: + problem = validate_routstr_auto_topup_settings(settings) + if problem is not None: logger.warning( - "Auto top-up enabled but missing configuration", - extra={ - "provider_id": row.id, - "has_threshold": bool(threshold), - "has_amount": bool(amount), - "has_mint": bool(mint_url), - }, + "Auto top-up enabled but its configuration is invalid", + extra={"provider_id": row.id, "problem": problem}, ) return + threshold_sats = _routstr_threshold_sats(row, settings) + amount = int(settings["topup_amount_limit"]) + mint_url = str(settings["topup_mint_url"]) + if not row.api_key: return @@ -235,16 +316,40 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: provider = RoutstrUpstreamProvider.from_db_row(row) if provider is None: return + if await _reconcile_routstr_state(row, provider): + return + balance = await provider.get_balance() - if balance is None: + if balance is None or not math.isfinite(balance) or balance < 0: logger.warning( "Could not fetch balance for auto top-up", extra={"provider_id": row.id, "base_url": row.base_url}, ) return - if balance >= threshold * 1000: + if balance >= threshold_sats: + return + + spent_24h_sats = await _routstr_spent_last_24h_sats() + if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS: + logger.critical( + "Auto top-up skipped: rolling 24h spend cap reached", + extra={ + "provider_id": row.id, + "spent_24h_sats": spent_24h_sats, + "topup_amount": amount, + "daily_cap_sats": ROUTSTR_MAX_DAILY_TOPUP_SATS, + }, + ) + return + + # The balance the peer must report before another token may be sent. Any + # shortfall is treated as "not credited": the token is a bearer instrument + # and a peer that took one without crediting it must not be handed another. + expected_sats = math.floor(balance) + amount + operation_id = await _claim_routstr_topup(row, expected_sats=expected_sats) + if operation_id is None: return # Balance is below threshold - create token and top up @@ -253,58 +358,64 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: extra={ "provider_id": row.id, "balance": balance, - "threshold": threshold, + "threshold_sats": threshold_sats, "topup_amount": amount, "mint_url": mint_url, }, ) try: - token = await send_token(amount, "sat", mint_url) + async with wallet_operation_guard(): + # Keep the spend cap and audit mutation in one wallet lock. + spent_24h_sats = await _routstr_spent_last_24h_sats() + if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS: + raise ValueError("Routstr auto top-up daily spend cap reached") + token = await send_token_from_owner_locked(amount, "sat", mint_url) + actual_mint_url = token_mint_url(token, mint_url) + try: + await _persist_routstr_token_and_mark_sent( + row, + operation_id, + expected_sats=expected_sats, + token=token, + amount=amount, + mint_url=actual_mint_url, + ) + except Exception: + logger.critical( + "Aborting auto top-up because its token and sent claim " + "could not be persisted atomically", + extra={"provider_id": row.id, "mint_url": actual_mint_url}, + ) + try: + await release_token_reservation(token) + except Exception as error: + logger.critical( + "Failed to release untracked auto-topup token", + extra={ + "provider_id": row.id, + "mint_url": actual_mint_url, + "error": repr(error), + }, + ) + else: + logger.warning( + "Auto-topup token was released after persistence failed", + extra={"provider_id": row.id, "mint_url": actual_mint_url}, + ) + raise except Exception as e: - logger.error( - "Failed to create cashu token for auto top-up", + logger.warning( + "Failed to create or persist cashu token for auto top-up", extra={ "provider_id": row.id, "amount": amount, "mint_url": mint_url, - "error": str(e), + "error": repr(e), + "error_type": type(e).__name__, }, ) - return - - actual_mint_url = token_mint_url(token, mint_url) - try: - await store_cashu_transaction( - token=token, - amount=amount, - unit="sat", - mint_url=actual_mint_url, - typ="out", - collected=False, - source="auto_topup", - ) - except Exception: - logger.critical( - "Aborting auto top-up because its cashu token could not be persisted", - extra={"provider_id": row.id, "mint_url": actual_mint_url}, - ) - try: - await release_token_reservation(token) - except Exception as error: - logger.critical( - "Failed to release untracked auto-topup token", - extra={ - "provider_id": row.id, - "mint_url": actual_mint_url, - "error": str(error), - }, - ) - else: - logger.warning( - "Auto-topup token was released after persistence failed", - extra={"provider_id": row.id, "mint_url": actual_mint_url}, - ) + await _release_routstr_claim(row, operation_id) return result = await provider.topup(token) @@ -348,6 +459,481 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: ) +def _routstr_state_id(row: UpstreamProviderRow) -> str: + if row.id is None: + raise ValueError("Routstr auto top-up requires a persisted provider row") + return _routstr_state_id_for_provider(row.id) + + +def _routstr_state_id_for_provider(provider_id: int | str) -> str: + return f"routstr-auto-topup-{provider_id}" + + +class RoutstrClaim(typing.NamedTuple): + operation_id: str + # Worker lease for "claimed"/"sent", retry-not-before for "backoff", and + # meaningless for "halted". + deadline: int + phase: str + # Peer balance in sats that proves this attempt was credited. + expected_sats: int + failures: int + + +def _routstr_request_id( + operation_id: str, + deadline: int, + phase: str, + expected_sats: int, + failures: int, +) -> str: + return f"routstr:{operation_id}:{deadline}:{phase}:{expected_sats}:{failures}" + + +def _parse_routstr_request_id(request_id: str | None) -> RoutstrClaim | None: + parts = (request_id or "").split(":", 5) + if len(parts) != 6 or parts[0] != "routstr" or parts[3] not in ROUTSTR_PHASES: + return None + try: + deadline = int(parts[2]) + expected_sats = int(parts[4]) + failures = int(parts[5]) + except (TypeError, ValueError): + return None + return RoutstrClaim(parts[1], deadline, parts[3], expected_sats, failures) + + +async def _routstr_spent_last_24h_sats() -> int: + """Total sats committed to Routstr auto top-ups in the last 24 hours. + + Uncollected rows count too: a token whose delivery is unconfirmed is spent + for capping purposes. Rows marked ``collected=False, swept=True`` record + tokens that were provably returned to the wallet and are excluded. + """ + cutoff = int(time.time()) - 24 * 60 * 60 + async with create_session() as session: + rows = ( + await session.exec( + select(CashuTransaction.amount, CashuTransaction.unit).where( + col(CashuTransaction.source) == "auto_topup", + col(CashuTransaction.type) == "out", + col(CashuTransaction.created_at) >= cutoff, + or_( + col(CashuTransaction.collected) == True, # noqa: E712 + col(CashuTransaction.swept) == False, # noqa: E712 + ), + ) + ) + ).all() + return sum( + amount if unit == "sat" else math.ceil(amount / 1000) for amount, unit in rows + ) + + +async def _routstr_provider_is_claimable( + session: AsyncSession, provider_id: int | str | None +) -> bool: + """Re-read the provider inside the claim transaction. + + Same reasoning as :func:`_ppq_provider_is_claimable`: a provider deleted or + retyped concurrently must either be visible here or lose the race against + the claim we are about to write. + """ + if provider_id is None: + return False + current = await session.get(UpstreamProviderRow, provider_id) + return current is not None and current.provider_type == "routstr" + + +async def _claim_routstr_topup( + row: UpstreamProviderRow, *, expected_sats: int +) -> str | None: + """Acquire the provider's single durable auto top-up slot. + + An expired backoff hands its failure count to the new attempt, so repeated + non-crediting peers still walk towards the halt instead of resetting the + counter every cycle. + """ + state_id = _routstr_state_id(row) + operation_id = uuid.uuid4().hex + deadline = int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS + + async with create_session() as session: + if not await _routstr_provider_is_claimable(session, row.id): + return None + existing = await session.get(CashuTransaction, state_id) + if existing is not None: + failures = 0 + if not (existing.collected or existing.swept): + claim = _parse_routstr_request_id(existing.request_id) + if ( + claim is None + or claim.phase != ROUTSTR_PHASE_BACKOFF + or time.time() < claim.deadline + ): + return None + failures = claim.failures + result = await session.exec( # type: ignore[call-overload] + update(CashuTransaction) + .where( + col(CashuTransaction.id) == state_id, + # Fence on the exact row that was read: any concurrent + # writer that moved the claim must win instead of us. + col(CashuTransaction.request_id) == existing.request_id, + ) + .values( + token="pending", + amount=0, + unit="sat", + mint_url=None, + request_id=_routstr_request_id( + operation_id, + deadline, + ROUTSTR_PHASE_CLAIMED, + expected_sats, + failures, + ), + collected=False, + swept=False, + created_at=int(time.time()), + source="routstr_auto_topup_claim", + ) + ) + await session.commit() + if (getattr(result, "rowcount", 0) or 0) != 1: + return None + return operation_id + + try: + async with create_session() as session: + if not await _routstr_provider_is_claimable(session, row.id): + return None + session.add( + CashuTransaction( + id=state_id, + token="pending", + amount=0, + unit="sat", + type="out", + request_id=_routstr_request_id( + operation_id, + deadline, + ROUTSTR_PHASE_CLAIMED, + expected_sats, + 0, + ), + collected=False, + source="routstr_auto_topup_claim", + ) + ) + await session.commit() + except IntegrityError: + return None + return operation_id + + +async def _advance_routstr_claim( + row: UpstreamProviderRow, + operation_id: str, + *, + deadline: int, + phase: str, + expected_sats: int, + failures: int, + token: str | None = None, + amount: int | None = None, + mint_url: str | None = None, +) -> bool: + """Move this worker's claim to another phase, if it still owns it.""" + values: dict[str, object] = { + "request_id": _routstr_request_id( + operation_id, deadline, phase, expected_sats, failures + ) + } + if token is not None: + values.update(token=token, amount=amount, mint_url=mint_url) + + async with create_session() as session: + result = await session.exec( # type: ignore[call-overload] + update(CashuTransaction) + .where( + col(CashuTransaction.id) == _routstr_state_id(row), + col(CashuTransaction.request_id).like(f"routstr:{operation_id}:%"), + col(CashuTransaction.collected) == False, # noqa: E712 + col(CashuTransaction.swept) == False, # noqa: E712 + ) + .values(**values) + ) + await session.commit() + return (getattr(result, "rowcount", 0) or 0) == 1 + + +async def _set_routstr_state_terminal( + row: UpstreamProviderRow, operation_id: str, *, collected: bool, swept: bool +) -> bool: + """Finish an attempt only if this worker still owns the claim.""" + async with create_session() as session: + result = await session.exec( # type: ignore[call-overload] + update(CashuTransaction) + .where( + col(CashuTransaction.id) == _routstr_state_id(row), + col(CashuTransaction.request_id).like(f"routstr:{operation_id}:%"), + col(CashuTransaction.collected) == False, # noqa: E712 + col(CashuTransaction.swept) == False, # noqa: E712 + ) + .values(collected=collected, swept=swept) + ) + await session.commit() + return (getattr(result, "rowcount", 0) or 0) == 1 + + +async def _release_routstr_claim(row: UpstreamProviderRow, operation_id: str) -> None: + """Hand back a claim whose token never left the wallet.""" + if not await _set_routstr_state_terminal( + row, operation_id, collected=False, swept=True + ): + logger.warning( + "Could not release the auto top-up claim after a pre-send failure; " + "it is owned by another attempt", + extra={"provider_id": row.id}, + ) + + +async def _persist_routstr_token_and_mark_sent( + row: UpstreamProviderRow, + operation_id: str, + *, + expected_sats: int, + token: str, + amount: int, + mint_url: str, +) -> None: + state_id = _routstr_state_id(row) + async with create_session() as session: + state = await session.get(CashuTransaction, state_id) + claim = _parse_routstr_request_id(state.request_id if state else None) + if ( + state is None + or state.collected + or state.swept + or claim is None + or claim.operation_id != operation_id + or claim.phase != ROUTSTR_PHASE_CLAIMED + ): + raise RuntimeError("Routstr auto top-up claim ownership was lost") + + result = await session.exec( # type: ignore[call-overload] + update(CashuTransaction) + .where( + col(CashuTransaction.id) == state_id, + col(CashuTransaction.request_id) == state.request_id, + col(CashuTransaction.collected) == False, # noqa: E712 + col(CashuTransaction.swept) == False, # noqa: E712 + ) + .values( + request_id=_routstr_request_id( + operation_id, + int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS, + ROUTSTR_PHASE_SENT, + expected_sats, + claim.failures, + ), + token=token, + amount=amount, + unit="sat", + mint_url=mint_url, + ) + ) + if (getattr(result, "rowcount", 0) or 0) != 1: + await session.rollback() + raise RuntimeError("Routstr auto top-up claim ownership was lost") + + session.add( + CashuTransaction( + id=uuid.uuid4().hex, + token=token, + amount=amount, + unit="sat", + mint_url=mint_url, + type="out", + collected=False, + source="auto_topup", + ) + ) + await session.commit() + + +async def _current_routstr_claim(row: UpstreamProviderRow) -> RoutstrClaim | None: + async with create_session() as session: + transaction = await session.get(CashuTransaction, _routstr_state_id(row)) + if transaction is None: + return None + return _parse_routstr_request_id(transaction.request_id) + + +async def _reconcile_routstr_state( + row: UpstreamProviderRow, provider: RoutstrUpstreamProvider +) -> bool: + """Return True while a prior attempt must suppress a new payment.""" + async with create_session() as session: + transaction = await session.get(CashuTransaction, _routstr_state_id(row)) + if transaction is None or transaction.collected or transaction.swept: + return False + + claim = _parse_routstr_request_id(transaction.request_id) + if claim is None: + logger.critical( + "Malformed auto top-up state; suppressing duplicate payment", + extra={"provider_id": row.id}, + ) + return True + + now = time.time() + if claim.phase == ROUTSTR_PHASE_HALTED: + return True + if claim.phase == ROUTSTR_PHASE_BACKOFF: + return now < claim.deadline + if claim.phase == ROUTSTR_PHASE_CLAIMED: + # Nothing left the wallet, so a dead worker's slot is free to reuse. + if now < claim.deadline: + return True + return not await _set_routstr_state_terminal( + row, claim.operation_id, collected=False, swept=True + ) + + balance = await provider.get_balance() + if ( + balance is not None + and math.isfinite(balance) + and balance >= claim.expected_sats + ): + if not await _set_routstr_state_terminal( + row, claim.operation_id, collected=True, swept=False + ): + logger.critical( + "Auto top-up was credited but its claim was already released; " + "a duplicate top-up is possible on the next cycle", + extra={"provider_id": row.id}, + ) + return True + if now < claim.deadline: + return True + + failures = claim.failures + 1 + if failures >= ROUTSTR_MAX_TOPUP_FAILURES: + await _advance_routstr_claim( + row, + claim.operation_id, + deadline=claim.deadline, + phase=ROUTSTR_PHASE_HALTED, + expected_sats=claim.expected_sats, + failures=failures, + ) + logger.critical( + "Auto top-up halted: the peer repeatedly failed to credit a token", + extra={ + "provider_id": row.id, + "base_url": row.base_url, + "failures": failures, + "admin_action": ( + f"POST /admin/api/upstream-providers/{row.id}" + "/routstr-auto-topup/release" + ), + }, + ) + return True + + await _advance_routstr_claim( + row, + claim.operation_id, + deadline=int(now) + ROUTSTR_BACKOFF_BASE_SECONDS * 2 ** (failures - 1), + phase=ROUTSTR_PHASE_BACKOFF, + expected_sats=claim.expected_sats, + failures=failures, + ) + logger.warning( + "Auto top-up was not credited by the peer; backing off", + extra={ + "provider_id": row.id, + "base_url": row.base_url, + "expected_sats": claim.expected_sats, + "failures": failures, + }, + ) + return True + + +async def get_routstr_auto_topup_state(provider_id: int) -> dict[str, object]: + """Return admin-safe state for a provider's durable Routstr claim.""" + async with create_session() as session: + transaction = await session.get( + CashuTransaction, _routstr_state_id_for_provider(provider_id) + ) + if transaction is None or transaction.collected or transaction.swept: + return {"active": False} + + claim = _parse_routstr_request_id(transaction.request_id) + return { + "active": True, + # Echoed back verbatim on release so a claim that moved on since the + # admin reviewed it fails the write instead of being swept unseen. + "state_token": transaction.request_id, + "operation_id": claim.operation_id if claim else None, + "phase": claim.phase if claim else None, + "expected_sats": claim.expected_sats if claim else None, + "failures": claim.failures if claim else None, + "deadline": claim.deadline if claim else None, + "created_at": transaction.created_at, + "amount": transaction.amount, + "unit": transaction.unit, + "mint_url": transaction.mint_url, + "malformed": claim is None, + } + + +class RoutstrReleaseOutcome(typing.NamedTuple): + released: bool + reason: str + + +async def release_routstr_auto_topup_state( + provider_id: int, *, state_token: str | None +) -> RoutstrReleaseOutcome: + """Clear a halted or stuck claim after an admin reconciles the peer. + + Unlike a Lightning melt there is no in-flight window to protect: the token + is already with the peer or still in the wallet either way. What the fence + does protect is the admin's decision — the row must be byte-identical to + the one they reviewed, so a claim that advanced in the meantime is not + swept on the strength of stale information. + """ + state_id = _routstr_state_id_for_provider(provider_id) + async with create_session() as session: + transaction = await session.get(CashuTransaction, state_id) + + if transaction is None or transaction.collected or transaction.swept: + return RoutstrReleaseOutcome(False, "no_active_claim") + if transaction.request_id != state_token: + return RoutstrReleaseOutcome(False, "stale_state") + + async with create_session() as session: + result = await session.exec( # type: ignore[call-overload] + update(CashuTransaction) + .where( + col(CashuTransaction.id) == state_id, + col(CashuTransaction.request_id) == state_token, + col(CashuTransaction.collected) == False, # noqa: E712 + col(CashuTransaction.swept) == False, # noqa: E712 + ) + .values(swept=True) + ) + if (getattr(result, "rowcount", 0) or 0) == 1: + await session.commit() + return RoutstrReleaseOutcome(True, "released") + await session.rollback() + return RoutstrReleaseOutcome(False, "claim_changed") + + def _ppq_state_id(row: UpstreamProviderRow) -> str: if row.id is None: raise ValueError("PPQ auto top-up requires a persisted provider row") @@ -528,7 +1114,11 @@ async def _set_ppq_state_terminal( col(CashuTransaction.collected) == False, # noqa: E712 col(CashuTransaction.swept) == False, # noqa: E712 ) - .values(collected=collected, swept=swept) + .values( + collected=collected, + swept=swept, + created_at=int(time.time()) if collected else CashuTransaction.created_at, + ) ) updated = (getattr(result, "rowcount", 0) or 0) == 1 if updated: @@ -553,8 +1143,10 @@ async def _reconcile_ppq_state( """ async with create_session() as session: transaction = await session.get(CashuTransaction, _ppq_state_id(row)) - if transaction is None or transaction.collected or transaction.swept: + if transaction is None or transaction.swept: return False + if transaction.collected: + return int(time.time()) - transaction.created_at < PPQ_SETTLED_COOLDOWN_SECONDS claim = _parse_ppq_request_id(transaction.request_id) if claim is None: @@ -620,7 +1212,7 @@ async def _reconcile_ppq_state( async def _ppq_provider_is_claimable( - session: AsyncSession, provider_id: int | None + session: AsyncSession, row: UpstreamProviderRow ) -> bool: """Re-read the provider inside the claim transaction. @@ -631,10 +1223,16 @@ async def _ppq_provider_is_claimable( this the worker could create a claim for a provider that no longer exists, orphaning it forever. """ - if provider_id is None: + if row.id is None: return False - current = await session.get(UpstreamProviderRow, provider_id) - return current is not None and current.provider_type == "ppqai" + current = await session.get(UpstreamProviderRow, row.id) + return bool( + current is not None + and current.enabled + and current.provider_type == "ppqai" + and current.api_key == row.api_key + and current.provider_settings == row.provider_settings + ) async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None: @@ -645,10 +1243,17 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None: request_id = _ppq_request_id(operation_id, expires_at, PPQ_PHASE_CLAIMED, "pending") async with create_session() as session: - if not await _ppq_provider_is_claimable(session, row.id): + if not await _ppq_provider_is_claimable(session, row): return None existing = await session.get(CashuTransaction, state_id) if existing is not None: + if ( + existing.collected + and not existing.swept + and int(time.time()) - existing.created_at + < PPQ_SETTLED_COOLDOWN_SECONDS + ): + return None result = await session.exec( # type: ignore[call-overload] update(CashuTransaction) .where( @@ -679,7 +1284,7 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None: async with create_session() as session: # Same fencing as the update path: the provider must still exist # inside the transaction that creates the claim. - if not await _ppq_provider_is_claimable(session, row.id): + if not await _ppq_provider_is_claimable(session, row): return None session.add( CashuTransaction( @@ -900,6 +1505,26 @@ async def _check_and_topup_ppq(row: UpstreamProviderRow, settings: dict) -> None if balance >= threshold_usd: return + # Require two low-balance reads before creating an invoice. + confirmed_balance = await provider.get_balance() + if ( + confirmed_balance is None + or not math.isfinite(confirmed_balance) + or confirmed_balance < 0 + or confirmed_balance >= threshold_usd + ): + logger.info( + "PPQ auto top-up aborted by balance confirmation", + extra={ + "provider_id": row.id, + "first_balance_usd": balance, + "confirmed_balance_usd": confirmed_balance, + "threshold_usd": threshold_usd, + }, + ) + return + balance = confirmed_balance + # Perform local pricing and owner-funds checks before asking PPQ to create # an invoice. The exact mint quote still has to be checked afterward, but # predictable local failures should not leave abandoned PPQ invoices. diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index bc9fd547..217cf9fc 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1,12 +1,13 @@ from __future__ import annotations import asyncio +import inspect import json import math import traceback import typing import uuid -from collections.abc import AsyncGenerator, AsyncIterator, Iterator +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterator from typing import Any, Mapping, Self, cast import httpx @@ -60,9 +61,10 @@ from .cache_breakpoints import ( inject_anthropic_cache_breakpoints, is_explicit_cache_model, ) -from .count_tokens import count_tokens_locally +from .count_tokens import MissingUsageEstimator, count_tokens_locally from .litellm_routing import detect_litellm_prefix from .rate_limit import UPSTREAM_RATE_LIMIT, classify_rate_limit +from .reasoning_effort import apply_reasoning_effort if typing.TYPE_CHECKING: from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget @@ -70,6 +72,32 @@ 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] @@ -83,6 +111,24 @@ def _cost_field( return value if isinstance(value, (int, float)) else default +def _settled_cost_msats(cost_data: CostMetadata) -> int: + charged = _cost_field(cost_data, "charged_msats", -1) + if charged >= 0: + return int(charged) + return int(_cost_field(cost_data, "total_msats")) + + +def _published_cost(cost_data: CostMetadata) -> dict[str, Any]: + cost = dict(cost_data) if isinstance(cost_data, dict) else cost_data.dict() + computed_msats = int(_cost_field(cost_data, "total_msats")) + settled_msats = _settled_cost_msats(cost_data) + if computed_msats != settled_msats: + cost["computed_msats"] = computed_msats + cost["total_msats"] = settled_msats + cost["charged_msats"] = settled_msats + return cost + + def _inject_cost_response_headers( headers: dict[str, str], cost_data: CostMetadata ) -> None: @@ -93,9 +139,11 @@ def _inject_cost_response_headers( usage tracking entry — without them, x-cashu requests show 0.0 for all sat cost fields. """ - headers["X-Routstr-Cost-Msats"] = str( - int(_cost_field(cost_data, "total_msats")) - ) + settled_msats = _settled_cost_msats(cost_data) + computed_msats = int(_cost_field(cost_data, "total_msats")) + headers["X-Routstr-Cost-Msats"] = str(settled_msats) + if computed_msats != settled_msats: + headers["X-Routstr-Computed-Cost-Msats"] = str(computed_msats) headers["X-Routstr-Input-Cost-Msats"] = str( int(_cost_field(cost_data, "input_msats")) ) @@ -107,6 +155,78 @@ def _inject_cost_response_headers( headers["X-Routstr-Cost-Usd"] = str(total_usd) +def _apply_estimated_usage( + response_json: dict[str, Any], + request_body: bytes | None, + model_obj: Model | None, + amount: int, + unit: str, + api: str, +) -> None: + """Bill a buffered response from a local estimate when usage is missing.""" + if response_json.get("usage"): + return + estimator = MissingUsageEstimator(request_body, model_obj) + estimator.observe(response_json) + estimated = estimator.estimated_usage(response_json.get("model")) + if not estimated: + return + logger.warning( + "No usage in non-streaming response, billing from local token estimate", + extra={ + "api": api, + "model": response_json.get("model", "unknown"), + "amount": amount, + "unit": unit, + "estimated_usage": estimated, + }, + ) + response_json["usage"] = estimated + + +def _parse_sse_events(content: str) -> list[tuple[list[str], str]]: + """Split a buffered SSE body into ``(field_lines, data)`` pairs. + + ``data`` is the newline-joined payload the SSE spec reassembles from every + ``data:`` line of one event, so multi-line JSON survives. Comment/keepalive + lines are dropped and events carrying no data at all are skipped; the + remaining ``event:``/``id:``/``retry:`` fields stay attached to their event + so Responses API framing is preserved on re-emission. A trailing event + without its blank-line terminator is still returned. + """ + events: list[tuple[list[str], str]] = [] + normalized = content.replace("\r\n", "\n").replace("\r", "\n") + for raw_event in normalized.split("\n\n"): + field_lines: list[str] = [] + data_lines: list[str] = [] + for line in raw_event.split("\n"): + if line.startswith("data:"): + data_lines.append(line[len("data:") :].lstrip(" ")) + elif line and not line.startswith(":"): + field_lines.append(line) + if not data_lines: + continue + events.append((field_lines, "\n".join(data_lines))) + return events + + +def _responses_usage_payload(data_json: dict) -> dict: + """Return the object carrying a Responses API event's model and usage. + + Canonical events nest them under ``response`` (``response.completed`` / + ``response.incomplete``); legacy and compat shapes keep them at top level. + """ + nested = data_json.get("response") + return nested if isinstance(nested, dict) else data_json + + +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) + body += "".join(f"data: {line}\n" for line in data.split("\n")) + return body + "\n" + + def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> None: """Inject cost breakdown into the response body's ``usage.cost`` object. @@ -122,11 +242,14 @@ def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> Non # data always overwrites any upstream-provided cost values. Using # setdefault would silently keep stale upstream values and drop our # calculated msats breakdown. + computed_msats = int(_cost_field(cost_data, "total_msats")) + settled_msats = _settled_cost_msats(cost_data) cost_obj: dict[str, int | float] = { "base_msats": int(_cost_field(cost_data, "base_msats")), "input_msats": int(_cost_field(cost_data, "input_msats")), "output_msats": int(_cost_field(cost_data, "output_msats")), - "total_msats": int(_cost_field(cost_data, "total_msats")), + "total_msats": settled_msats, + "charged_msats": settled_msats, "cache_read_input_tokens": int( _cost_field(cost_data, "cache_read_input_tokens") ), @@ -134,15 +257,15 @@ def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> Non _cost_field(cost_data, "cache_creation_input_tokens") ), "cache_read_msats": int(_cost_field(cost_data, "cache_read_msats")), - "cache_creation_msats": int( - _cost_field(cost_data, "cache_creation_msats") - ), + "cache_creation_msats": int(_cost_field(cost_data, "cache_creation_msats")), } + if computed_msats != settled_msats: + cost_obj["computed_msats"] = computed_msats total_usd = float(_cost_field(cost_data, "total_usd", 0.0)) if total_usd: cost_obj["total_usd"] = total_usd usage["cost"] = cost_obj - usage["cost_sats"] = int(_cost_field(cost_data, "total_msats")) // 1000 + usage["cost_sats"] = settled_msats // 1000 def _is_json_content_type(content_type: str | None) -> bool: @@ -155,6 +278,35 @@ def _is_json_content_type(content_type: str | None) -> bool: return main.startswith("application/") and main.endswith("+json") +def _is_sse_body(content_type: str | None, content_str: str) -> bool: + if content_type: + main = content_type.split(";", 1)[0].strip().lower() + if main == "text/event-stream": + return True + if _is_json_content_type(content_type): + return False + + for line in content_str.lstrip("\ufeff").splitlines(): + stripped = line.strip() + if stripped: + return stripped.startswith(("data:", "event:", "id:", "retry:", ":")) + return False + + +def _openai_completion_path(path: str) -> str | None: + canonical = "/" + path.rstrip("/") + if canonical.endswith("/chat/completions"): + return "chat/completions" + return "completions" if canonical.endswith("/completions") else None + + +def _x_cashu_path_has_settlement_handler(path: str) -> bool: + canonical = path.rstrip("/") + return _openai_completion_path(canonical) is not None or canonical.endswith( + ("embeddings", "messages", "messages/count_tokens") + ) + + class TopupData(BaseModel): """Universal top-up data schema for Lightning Network invoices.""" @@ -341,6 +493,34 @@ class BaseUpstreamProvider: return response_json["provider"] = f"{provider_type}:{existing_str}" + def _log_full_refund( + self, + *, + route: str, + model: str | None, + content_str: str, + amount: int, + unit: str, + ) -> None: + """Record a settlement that serves content but charges nothing. + + The client keeps both the response and the whole prepayment, so the + model, the serving upstream and a redacted body preview are logged to + keep the unbilled request auditable. + """ + logger.warning( + "Zero-cost settlement, refunding the full prepayment", + extra={ + "route": route, + "model": model or "unknown", + "provider_type": self.provider_type, + "upstream_base_url": self.base_url, + "refund_amount": amount, + "unit": unit, + "response_body_preview": redact_org_ids(content_str.strip()[:500]), + }, + ) + def inject_cost_metadata( self, response_json: dict, @@ -349,14 +529,8 @@ class BaseUpstreamProvider: ) -> None: """Unifies the injection of cost and usage metadata across all completion types.""" self._apply_provider_field(response_json) - if isinstance(cost_data, dict): - total_msats = cost_data.get("total_msats", 0) - cost_dict = cost_data - else: - total_msats = cost_data.total_msats - cost_dict = cost_data.dict() - - sats_cost = total_msats // 1000 + cost_dict = _published_cost(cost_data) + sats_cost = cost_dict["total_msats"] // 1000 # Inject the shared SDK cost contract into every usage shape. if isinstance(response_json.get("usage"), dict): @@ -370,6 +544,14 @@ class BaseUpstreamProvider: message["usage"]["remaining_balance_msats"] = key.balance self._fold_cache_into_input_tokens(message["usage"]) + nested_response = response_json.get("response") + if isinstance(nested_response, dict) and isinstance( + nested_response.get("usage"), dict + ): + _inject_cost_into_usage(nested_response, cost_data) + nested_response["usage"]["remaining_balance_msats"] = key.balance + self._fold_cache_into_input_tokens(nested_response["usage"]) + # Unified Routstr metadata response_json["metadata"] = response_json.get("metadata", {}) response_json["metadata"]["routstr"] = { @@ -409,6 +591,7 @@ class BaseUpstreamProvider: "refund-lnurl", "key-expiry-time", "x-cashu", + "x-routstr-model-path", ]: if headers.pop(header, None) is not None: removed_headers.append(header) @@ -529,8 +712,7 @@ class BaseUpstreamProvider: transformed_model = self.transform_model_name(original_model) data["input"]["model"] = transformed_model - # Ensure proper Responses API structure - # Add any Responses-specific transformations here + apply_reasoning_effort(data, model_obj) return json.dumps(data).encode() except Exception as e: @@ -558,17 +740,20 @@ class BaseUpstreamProvider: return "openrouter.ai" in (self.base_url or "") def prepare_request_body( - self, body: bytes | None, model_obj: Model + self, + body: bytes | None, + model_obj: Model, + include_stream_usage: bool = False, ) -> bytes | None: """Transform request body for provider-specific requirements. - Automatically transforms model names and, for streaming chat - completions, opts the upstream into emitting per-chunk ``usage`` - so cost tracking can read real token counts instead of falling - back to ``MaxCostData``. + Automatically transforms model names and opts streaming OpenAI + completion endpoints into emitting per-chunk ``usage`` so cost + tracking can read real token counts. Args: body: Original request body bytes + include_stream_usage: Opt a streaming completion into usage chunks Returns: Transformed request body bytes @@ -610,15 +795,9 @@ class BaseUpstreamProvider: # OpenAI-compatible streaming responses omit ``usage`` unless the # request sets ``stream_options.include_usage = true``. Without it - # we can't reconcile token counts at end of stream and the - # request gets billed at max-cost with zero tokens. Discriminate - # chat-completions-shaped requests by the ``messages`` field so we - # don't poke unrelated endpoints. - if ( - data.get("stream") is True - and "messages" in data - and isinstance(data.get("messages"), list) - ): + # we can't reconcile token counts at end of stream and must use + # the local request/response estimator. + if data.get("stream") is True and include_stream_usage: existing = data.get("stream_options") merged = dict(existing) if isinstance(existing, dict) else {} if merged.get("include_usage") is not True: @@ -647,6 +826,9 @@ class BaseUpstreamProvider: if inject_anthropic_cache_breakpoints(data): changed = True + if apply_reasoning_effort(data, model_obj): + changed = True + if changed: return json.dumps(data).encode() return body @@ -925,6 +1107,9 @@ 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, + legacy_completion: bool = False, ) -> StreamingResponse: """Handle streaming chat completion responses with token usage tracking and cost adjustment. @@ -945,6 +1130,8 @@ class BaseUpstreamProvider: snapshot_key, snapshot_session ) + usage_estimator = MissingUsageEstimator(request_body, model_obj) + logger.debug( "Processing streaming chat completion", extra={ @@ -961,28 +1148,43 @@ class BaseUpstreamProvider: last_model_seen: str | None = None 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 - 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, - {"model": last_model_seen or "unknown", "usage": None}, - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - except Exception: - pass + 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 @@ -1005,7 +1207,7 @@ class BaseUpstreamProvider: * ``[DONE]`` is swallowed so it can be re-emitted exactly once at end of stream. """ - nonlocal last_model_seen, usage_chunk_data, done_seen + nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id event = raw_event.strip(b"\r\n") if not event: @@ -1047,6 +1249,7 @@ class BaseUpstreamProvider: obj = None if isinstance(obj, dict): + usage_estimator.observe(obj) self._apply_provider_field(obj) if obj.get("model"): last_model_seen = str(obj.get("model")) @@ -1057,9 +1260,12 @@ class BaseUpstreamProvider: or not isinstance(obj["id"], str) or obj["id"] == "existing-id" ): - if not hasattr(self, "_current_stream_id"): - self._current_stream_id = f"chatcmpl-{uuid.uuid4()}" - obj["id"] = self._current_stream_id + if stream_id is None: + id_prefix = "cmpl" if legacy_completion else "chatcmpl" + stream_id = f"{id_prefix}-{uuid.uuid4()}" + obj["id"] = stream_id + else: + stream_id = obj["id"] if isinstance(obj.get("usage"), dict): # Capture usage for end-of-stream cost reconciliation. # Some models (e.g. Gemini thinking models over the @@ -1136,13 +1342,8 @@ class BaseUpstreamProvider: if fresh_key: cost_data: dict try: - adjustment_input = ( - usage_chunk_data - if usage_chunk_data is not None - else { - "model": last_model_seen or "unknown", - "usage": None, - } + adjustment_input = usage_estimator.billing_data( + usage_chunk_data, last_model_seen ) cost_data = await adjust_payment_for_tokens( fresh_key, @@ -1176,11 +1377,14 @@ class BaseUpstreamProvider: raise if usage_chunk_data is None: - if not hasattr(self, "_current_stream_id"): - self._current_stream_id = f"chatcmpl-{uuid.uuid4()}" + if stream_id is None: + id_prefix = "cmpl" if legacy_completion else "chatcmpl" + stream_id = f"{id_prefix}-{uuid.uuid4()}" usage_chunk_data = { - "id": self._current_stream_id, - "object": "chat.completion.chunk", + "id": stream_id, + "object": "text_completion" + if legacy_completion + else "chat.completion.chunk", "model": last_model_seen or "unknown", "choices": [], "usage": { @@ -1212,18 +1416,24 @@ class BaseUpstreamProvider: except Exception as stream_error: logger.warning( - "Streaming interrupted; finalizing in background", + "Streaming interrupted; finalizing before closing upstream", extra={ "error": str(stream_error), + "error_type": type(stream_error).__name__, "key_hash": key.hashed_key[:8] + "...", }, ) raise finally: - if not usage_finalized: - # Create a background task to ensure finalization happens - # even if the generator is closed early - background_tasks.add_task(finalize_db_only) + # 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, + ) + ) # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -1245,6 +1455,8 @@ class BaseUpstreamProvider: requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, + request_body: bytes | None = None, + legacy_completion: bool = False, ) -> Response: """Handle non-streaming chat completion responses with token usage tracking and cost adjustment. @@ -1284,7 +1496,16 @@ class BaseUpstreamProvider: if requested_model: response_json["model"] = requested_model if "id" not in response_json or not isinstance(response_json["id"], str): - response_json["id"] = f"chatcmpl-{uuid.uuid4()}" + prefix = "cmpl" if legacy_completion else "chatcmpl" + response_json["id"] = f"{prefix}-{uuid.uuid4()}" + + usage = response_json.get("usage") + if not isinstance(usage, dict) or not usage: + usage_estimator = MissingUsageEstimator(request_body, model_obj) + usage_estimator.observe(response_json) + response_json["usage"] = usage_estimator.openai_response_data( + response_json.get("model") + )["usage"] cost_data = await adjust_payment_for_tokens( key, @@ -1307,18 +1528,12 @@ class BaseUpstreamProvider: ) self._fold_cache_into_input_tokens(response_json["usage"]) - # Keep detailed cost + published_cost = _published_cost(cost_data) + published_cost["sats_cost"] = published_cost["total_msats"] // 1000 + published_cost["remaining_balance_msats"] = remaining_balance_msats response_json["metadata"] = response_json.get("metadata", {}) - response_json["metadata"]["routstr"] = {"cost": cost_data} - response_json["metadata"]["routstr"]["cost"]["sats_cost"] = ( - cost_data.get("total_msats", 0) // 1000 - ) - response_json["metadata"]["routstr"]["cost"]["remaining_balance_msats"] = ( - remaining_balance_msats - ) - response_json["cost"] = cost_data - response_json["cost"]["sats_cost"] = cost_data.get("total_msats", 0) // 1000 - response_json["cost"]["remaining_balance_msats"] = remaining_balance_msats + response_json["metadata"]["routstr"] = {"cost": published_cost.copy()} + response_json["cost"] = published_cost logger.debug( "Payment adjustment completed for non-streaming", @@ -1389,6 +1604,8 @@ 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. @@ -1400,6 +1617,8 @@ class BaseUpstreamProvider: Returns: StreamingResponse with cost data injected at the end """ + usage_estimator = MissingUsageEstimator(request_body, model_obj) + logger.debug( "Processing streaming Responses API completion", extra={ @@ -1422,23 +1641,37 @@ class BaseUpstreamProvider: nonlocal usage_finalized if usage_finalized: return - 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, - {"model": last_model_seen or "unknown", "usage": None}, - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - except Exception: - pass + 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 @@ -1508,8 +1741,11 @@ class BaseUpstreamProvider: "response.incomplete", ): usage_chunk_data = obj + if not usage_estimator.output_text: + usage_estimator.observe(obj) return + usage_estimator.observe(obj) yield prefix + b"data: " + json.dumps(obj).encode() + b"\n\n" else: if final: @@ -1549,13 +1785,8 @@ class BaseUpstreamProvider: if fresh_key: cost_data: dict try: - adjustment_input = ( - usage_chunk_data - if usage_chunk_data is not None - else { - "model": last_model_seen or "unknown", - "usage": None, - } + adjustment_input = usage_estimator.billing_data( + usage_chunk_data, last_model_seen ) cost_data = await adjust_payment_for_tokens( fresh_key, @@ -1609,24 +1840,6 @@ class BaseUpstreamProvider: }, } - remaining_balance_msats = fresh_key.balance - sats_cost = cost_data.get("total_msats", 0) // 1000 - - if ( - "response" in usage_chunk_data - and isinstance(usage_chunk_data["response"], dict) - and "usage" in usage_chunk_data["response"] - ): - usage_chunk_data["response"]["usage"]["cost"] = ( - cost_data.get("total_usd", 0.0) - ) - usage_chunk_data["response"]["usage"]["cost_sats"] = ( - sats_cost - ) - usage_chunk_data["response"]["usage"][ - "remaining_balance_msats" - ] = remaining_balance_msats - try: self.inject_cost_metadata( usage_chunk_data, cost_data, fresh_key @@ -1646,16 +1859,24 @@ class BaseUpstreamProvider: except Exception as stream_error: logger.warning( - "Responses API streaming interrupted; finalizing in background", + "Responses API streaming interrupted; finalizing before closing upstream", extra={ "error": str(stream_error), + "error_type": type(stream_error).__name__, "key_hash": key.hashed_key[:8] + "...", }, ) raise finally: - if not usage_finalized: - await finalize_db_only() + # 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, + ) + ) # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -1677,6 +1898,7 @@ class BaseUpstreamProvider: requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, + request_body: bytes | None = None, ) -> Response: """Handle non-streaming Responses API responses with token usage tracking and cost adjustment. @@ -1716,6 +1938,13 @@ class BaseUpstreamProvider: }, ) + if not isinstance(response_json.get("usage"), dict): + usage_estimator = MissingUsageEstimator(request_body, model_obj) + usage_estimator.observe(response_json) + response_json["usage"] = usage_estimator.response_data( + response_json.get("model") + )["usage"] + if requested_model: response_json["model"] = requested_model if "id" not in response_json or not isinstance(response_json["id"], str): @@ -1742,18 +1971,12 @@ class BaseUpstreamProvider: ) self._fold_cache_into_input_tokens(response_json["usage"]) - # Keep detailed cost + published_cost = _published_cost(cost_data) + published_cost["sats_cost"] = published_cost["total_msats"] // 1000 + published_cost["remaining_balance_msats"] = remaining_balance_msats response_json["metadata"] = response_json.get("metadata", {}) - response_json["metadata"]["routstr"] = {"cost": cost_data} - response_json["metadata"]["routstr"]["cost"]["sats_cost"] = ( - cost_data.get("total_msats", 0) // 1000 - ) - response_json["metadata"]["routstr"]["cost"]["remaining_balance_msats"] = ( - remaining_balance_msats - ) - response_json["cost"] = cost_data - response_json["cost"]["sats_cost"] = cost_data.get("total_msats", 0) // 1000 - response_json["cost"]["remaining_balance_msats"] = remaining_balance_msats + response_json["metadata"]["routstr"] = {"cost": published_cost.copy()} + response_json["cost"] = published_cost logger.debug( "Payment adjustment completed for non-streaming Responses API", @@ -1836,9 +2059,9 @@ class BaseUpstreamProvider: return try: - # Finalize with "unknown" model and no usage to release reservation/charge max cost - # (no routed identity here by design: the None usage settles at - # MaxCostData before any pricing lookup can happen). + # Generic opaque streams have no request/response token seam. + # Missing usage therefore releases the reservation; the hold is + # never treated as evidence of consumption. await adjust_payment_for_tokens( key, {"model": "unknown", "usage": None}, @@ -1873,7 +2096,10 @@ class BaseUpstreamProvider: requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, + request_body: bytes | None = None, ) -> StreamingResponse: + usage_estimator = MissingUsageEstimator(request_body, model_obj) + async def stream_with_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: @@ -1927,13 +2153,9 @@ class BaseUpstreamProvider: usage_finalized = True return None try: - fallback: dict = { - "model": last_model_seen or "unknown", - "usage": None, - } cost_data = await adjust_payment_for_tokens( fresh_key, - fallback, + usage_estimator.response_data(last_model_seen), new_session, max_cost_for_model, model_obj, @@ -1972,6 +2194,7 @@ class BaseUpstreamProvider: try: data = json.loads(line[6:]) if isinstance(data, dict): + usage_estimator.observe(data) msg = data.get("message", {}) if msg and msg.get("model"): last_model_seen = str(msg.get("model")) @@ -2169,6 +2392,7 @@ class BaseUpstreamProvider: requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, + request_body: bytes | None = None, ) -> Response: try: content = await response.aread() @@ -2187,6 +2411,12 @@ class BaseUpstreamProvider: if path.endswith("count_tokens") and "usage" not in response_json: input_tokens = response_json.get("input_tokens", 0) response_json["usage"] = {"input_tokens": input_tokens} + elif not isinstance(response_json.get("usage"), dict): + usage_estimator = MissingUsageEstimator(request_body, model_obj) + usage_estimator.observe(response_json) + response_json["usage"] = usage_estimator.response_data( + response_json.get("model") + )["usage"] cost_data = await adjust_payment_for_tokens( key, @@ -2294,11 +2524,18 @@ class BaseUpstreamProvider: requested_model, model_obj, reservation_snapshot, + request_body, ) response_json = messages_dispatch.coerce_litellm_payload(result) if requested_model and "model" in response_json: response_json["model"] = requested_model + if not isinstance(response_json.get("usage"), dict): + usage_estimator = MissingUsageEstimator(request_body, model_obj) + usage_estimator.observe(response_json) + response_json["usage"] = usage_estimator.response_data( + response_json.get("model") + )["usage"] cost_data = await adjust_payment_for_tokens( key, @@ -2412,10 +2649,13 @@ class BaseUpstreamProvider: requested_model: str | None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, + request_body: bytes | None = None, ) -> StreamingResponse: """Re-emit a litellm Anthropic-event iterator as live SSE bytes with cost reconciliation appended at end of stream.""" + usage_estimator = MissingUsageEstimator(request_body, model_obj) + async def stream_with_cost() -> AsyncGenerator[bytes, None]: usage_finalized = False last_model_seen: str | None = None @@ -2432,12 +2672,10 @@ class BaseUpstreamProvider: if usage_finalized: return None logger.warning( - "Finalizing /v1/messages stream with no usage data — " - "client will be billed at max-cost with zero tokens. " - "Likely cause: upstream omitted `usage` from the SSE " - "stream (check that the request includes " - "`stream_options.include_usage=true` and that the " - "upstream actually emits a final usage chunk).", + "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", @@ -2451,13 +2689,9 @@ class BaseUpstreamProvider: usage_finalized = True return None try: - fallback: dict = { - "model": last_model_seen or "unknown", - "usage": None, - } cost_data = await adjust_payment_for_tokens( fresh_key, - fallback, + usage_estimator.response_data(last_model_seen), new_session, max_cost_for_model, model_obj, @@ -2490,6 +2724,7 @@ class BaseUpstreamProvider: async for annotated in messages_dispatch.stream_annotated_events( iterator, requested_model ): + usage_estimator.observe(annotated.event) if annotated.model: last_model_seen = annotated.model # Anthropic SSE reports usage cumulatively across @@ -2747,9 +2982,7 @@ class BaseUpstreamProvider: event_type = str(event.get("type") or "") prefix = f"event: {event_type}\n" if event_type else "" buffered[index] = annotated._replace( - sse_bytes=( - f"{prefix}data: {json.dumps(event)}\n\n".encode() - ) + sse_bytes=(f"{prefix}data: {json.dumps(event)}\n\n".encode()) ) async def replay() -> AsyncGenerator[bytes, None]: @@ -2788,6 +3021,7 @@ class BaseUpstreamProvider: Returns: Response or StreamingResponse from upstream with cost tracking """ + completion_path = _openai_completion_path(path) path = self.normalize_request_path(path, model_obj) if ( @@ -2816,7 +3050,11 @@ class BaseUpstreamProvider: (model_obj.forwarded_model_id or model_obj.id) if model_obj else None ) - transformed_body = self.prepare_request_body(request_body, model_obj) + transformed_body = self.prepare_request_body( + request_body, + model_obj, + include_stream_usage=completion_path is not None, + ) logger.debug( "Forwarding request to upstream", @@ -2913,7 +3151,7 @@ class BaseUpstreamProvider: return mapped_error if ( - path.endswith("chat/completions") + completion_path is not None or path.endswith("embeddings") or path.endswith("messages") or path.endswith("messages/count_tokens") @@ -2939,6 +3177,7 @@ class BaseUpstreamProvider: requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + request_body=request_body, ) background_tasks = BackgroundTasks() background_tasks.add_task(response.aclose) @@ -2957,6 +3196,7 @@ class BaseUpstreamProvider: requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + request_body=request_body, ) finally: await response.aclose() @@ -2974,12 +3214,13 @@ class BaseUpstreamProvider: requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + request_body=request_body, ) finally: await response.aclose() await client.aclose() - if path.endswith("chat/completions"): + if completion_path is not None: client_wants_streaming = False if request_body: try: @@ -3015,9 +3256,7 @@ class BaseUpstreamProvider: if is_streaming and response.status_code == 200: background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - result = await self.handle_streaming_chat_completion( + return await self.handle_streaming_chat_completion( response, key, max_cost_for_model, @@ -3025,9 +3264,10 @@ class BaseUpstreamProvider: requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + client=client, + request_body=request_body, + legacy_completion=completion_path == "completions", ) - result.background = background_tasks - return result # Handle both non-streaming chat completions and embeddings if response.status_code == 200: @@ -3040,6 +3280,8 @@ class BaseUpstreamProvider: requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + request_body=request_body, + legacy_completion=completion_path == "completions", ) finally: await response.aclose() @@ -3296,19 +3538,16 @@ class BaseUpstreamProvider: ) if is_streaming and response.status_code == 200: - result = await self.handle_streaming_responses_completion( + return 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, ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - result.background = background_tasks - return result if response.status_code == 200: try: @@ -3320,6 +3559,7 @@ class BaseUpstreamProvider: requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + request_body=transformed_body, ) finally: await response.aclose() @@ -3549,23 +3789,17 @@ class BaseUpstreamProvider: ) return cost case CostDataError() as error: + # Content was already served, so refund instead of raising. logger.error( - "Cost calculation error", + "Cost calculation error, refunding the prepayment", extra={ "model": model, "error_message": error.message, "error_code": error.code, }, ) - raise HTTPException( - status_code=400, - detail={ - "error": { - "message": error.message, - "type": "invalid_request_error", - "code": error.code, - } - }, + return MaxCostData( + base_msats=0, input_msats=0, output_msats=0, total_msats=0 ) return None @@ -3592,54 +3826,30 @@ class BaseUpstreamProvider: extra={"amount": amount, "unit": unit, "mint": mint}, ) - max_retries = 3 - last_exception = None - refund_token = None - - for attempt in range(max_retries): - try: - refund_token = await send_token(amount, unit=unit, mint_url=mint) - break - except Exception as e: - last_exception = e - if attempt < max_retries - 1: - logger.warning( - "Refund token creation failed, retrying", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "attempt": attempt + 1, - "max_retries": max_retries, - "amount": amount, - "unit": unit, - "mint": mint, - }, - ) - else: - logger.error( - "Failed to create refund token after all retries", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "attempt": attempt + 1, - "max_retries": max_retries, - "amount": amount, - "unit": unit, - "mint": mint, - }, - ) - - if refund_token is None: + try: + # Token creation may swap proofs, so it is unsafe to retry. + refund_token = await send_token(amount, unit=unit, mint_url=mint) + except Exception as error: + logger.error( + "Failed to create refund token", + extra={ + "error": str(error), + "error_type": type(error).__name__, + "amount": amount, + "unit": unit, + "mint": mint, + }, + ) raise HTTPException( status_code=401, detail={ "error": { - "message": f"failed to create refund after {max_retries} attempts: {str(last_exception)}", + "message": f"failed to create refund: {error}", "type": "invalid_request_error", "code": "send_token_failed", } }, - ) + ) from error logger.info( "Refund token created successfully", @@ -3647,7 +3857,6 @@ class BaseUpstreamProvider: "amount": amount, "unit": unit, "mint": mint, - "attempt": attempt + 1, "token_preview": refund_token[:20] + "..." if len(refund_token) > 20 else refund_token, @@ -3673,6 +3882,7 @@ class BaseUpstreamProvider: mint: str | None = None, request_id: str | None = None, model_obj: Model | None = None, + request_body: bytes | None = None, ) -> StreamingResponse: """Handle streaming response for X-Cashu payment, calculating refund if needed. @@ -3704,6 +3914,7 @@ class BaseUpstreamProvider: usage_data = None model = None cost_data: CostData | MaxCostData | None = None + usage_estimator = MissingUsageEstimator(request_body, model_obj) lines = content_str.strip().split("\n") for line in lines: @@ -3731,102 +3942,118 @@ class BaseUpstreamProvider: usage_data = merged except json.JSONDecodeError: continue + usage_estimator.observe(data_json) - if usage_data and model: - logger.debug( - "Found usage data in streaming response", + if not usage_data: + usage_data = usage_estimator.estimated_usage(model) + if usage_data: + logger.warning( + "No usage in streaming response, billing from local token estimate", + extra={ + "model": model, + "amount": amount, + "unit": unit, + "estimated_usage": usage_data, + }, + ) + + logger.debug( + "Calculating cost for streaming response", + extra={ + "model": model, + "usage_data": usage_data, + "amount": amount, + "unit": unit, + }, + ) + + response_data = {"usage": usage_data, "model": model or "unknown"} + try: + cost_data = await self.get_x_cashu_cost( + response_data, max_cost_for_model, model_obj + ) + if cost_data is not None and cost_data.total_msats == 0: + self._log_full_refund( + route="chat.streaming", + model=model, + content_str=content_str, + amount=amount, + unit=unit, + ) + if cost_data: + if unit == "msat": + refund_amount = amount - cost_data.total_msats + elif unit == "sat": + refund_amount = amount - (cost_data.total_msats + 999) // 1000 + else: + raise ValueError(f"Invalid unit: {unit}") + + if refund_amount > 0: + logger.debug( + "Processing refund for streaming response", + extra={ + "original_amount": amount, + "cost_msats": cost_data.total_msats, + "refund_amount": refund_amount, + "unit": unit, + "model": model, + }, + ) + + refund_token = await self.send_refund( + refund_amount, + unit, + mint, + request_id=request_id, + ) + response_headers["X-Cashu"] = refund_token + + logger.info( + "Refund processed for streaming response", + extra={ + "refund_amount": refund_amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, + }, + ) + else: + logger.debug( + "No refund needed for streaming response", + extra={ + "amount": amount, + "cost_msats": cost_data.total_msats, + "model": model, + }, + ) + + # Inject cost breakdown headers so the SDK's + # extractUsageFromResponseHeaders can populate + # inputMsats/outputMsats/totalMsats for x-cashu requests. + _inject_cost_response_headers(response_headers, cost_data) + except Exception as e: + logger.error( + "Error calculating cost for streaming response", extra={ + "error": str(e), + "error_type": type(e).__name__, "model": model, - "usage_data": usage_data, "amount": amount, "unit": unit, }, ) - response_data = {"usage": usage_data, "model": model} - try: - cost_data = await self.get_x_cashu_cost( - response_data, max_cost_for_model, model_obj - ) - if cost_data: - if unit == "msat": - refund_amount = amount - cost_data.total_msats - elif unit == "sat": - refund_amount = amount - (cost_data.total_msats + 999) // 1000 - else: - raise ValueError(f"Invalid unit: {unit}") - - if refund_amount > 0: - logger.debug( - "Processing refund for streaming response", - extra={ - "original_amount": amount, - "cost_msats": cost_data.total_msats, - "refund_amount": refund_amount, - "unit": unit, - "model": model, - }, - ) - - refund_token = await self.send_refund( - refund_amount, - unit, - mint, - request_id=request_id, - ) - response_headers["X-Cashu"] = refund_token - - logger.info( - "Refund processed for streaming response", - extra={ - "refund_amount": refund_amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - else: - logger.debug( - "No refund needed for streaming response", - extra={ - "amount": amount, - "cost_msats": cost_data.total_msats, - "model": model, - }, - ) - - # Inject cost breakdown headers so the SDK's - # extractUsageFromResponseHeaders can populate - # inputMsats/outputMsats/totalMsats for x-cashu requests. - _inject_cost_response_headers(response_headers, cost_data) - except Exception as e: - logger.error( - "Error calculating cost for streaming response", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "model": model, - "amount": amount, - "unit": unit, - }, - ) - for i, line in enumerate(lines): if line.startswith("data: "): try: data_json = json.loads(line[6:]) if not isinstance(data_json, dict): continue - changed = False - if "provider" not in data_json: - self._apply_provider_field(data_json) - changed = True - if ( - cost_data - and "usage" in data_json - and data_json["usage"] - ): + provider_before = data_json.get("provider") + self._apply_provider_field(data_json) + 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) changed = True if changed: @@ -3855,6 +4082,7 @@ class BaseUpstreamProvider: mint: str | None = None, request_id: str | None = None, model_obj: Model | None = None, + request_body: bytes | None = None, ) -> Response: """Handle non-streaming response for X-Cashu payment, calculating refund if needed. @@ -3876,9 +4104,20 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) self._apply_provider_field(response_json) + _apply_estimated_usage( + response_json, request_body, model_obj, amount, unit, "chat" + ) cost_data = await self.get_x_cashu_cost( response_json, max_cost_for_model, model_obj ) + if cost_data is not None and cost_data.total_msats == 0: + self._log_full_refund( + route="chat", + model=response_json.get("model"), + content_str=content_str, + amount=amount, + unit=unit, + ) if cost_data and "usage" in response_json: # Inject cost breakdown into both the response body (so the @@ -4010,6 +4249,7 @@ class BaseUpstreamProvider: mint: str | None = None, request_id: str | None = None, model_obj: Model | None = None, + request_body: bytes | None = None, ) -> StreamingResponse | Response: """Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming. @@ -4032,7 +4272,9 @@ class BaseUpstreamProvider: content_str = ( content.decode("utf-8") if isinstance(content, bytes) else content ) - is_streaming = content_str.startswith("data:") or "data:" in content_str + is_streaming = _is_sse_body( + response.headers.get("content-type"), content_str + ) logger.debug( "Chat completion response analysis", @@ -4054,6 +4296,7 @@ class BaseUpstreamProvider: mint, request_id=request_id, model_obj=model_obj, + request_body=request_body, ) else: return await self.handle_x_cashu_non_streaming_response( @@ -4065,6 +4308,7 @@ class BaseUpstreamProvider: mint, request_id=request_id, model_obj=model_obj, + request_body=request_body, ) except Exception as e: @@ -4093,6 +4337,8 @@ class BaseUpstreamProvider: max_cost_for_model: int, model_obj: Model, mint: str | None = None, + *, + request_body: bytes | None = None, ) -> Response | StreamingResponse: """Forward request paid with X-Cashu token to upstream service. @@ -4108,16 +4354,26 @@ class BaseUpstreamProvider: Returns: Response or StreamingResponse with refund if applicable """ + completion_path = _openai_completion_path(path) if path.startswith("v1/"): path = path.replace("v1/", "") - request_body = await request.body() + if request_body is None: + request_body = await request.body() if ( path.endswith("messages/count_tokens") and not self.supports_anthropic_messages ): - return count_tokens_locally(request_body, model_obj) + result = count_tokens_locally(request_body, model_obj) + refund_token = await self.send_refund( + amount, + unit, + mint, + request_id=getattr(request.state, "request_id", None), + ) + result.headers["X-Cashu"] = refund_token + return result if ( path.endswith("messages") @@ -4136,7 +4392,11 @@ class BaseUpstreamProvider: url = f"{self.base_url}/{path}" - transformed_body = self.prepare_request_body(request_body, model_obj) + transformed_body = self.prepare_request_body( + request_body, + model_obj, + include_stream_usage=completion_path is not None, + ) logger.debug( "Forwarding request to upstream", @@ -4232,12 +4492,7 @@ class BaseUpstreamProvider: error_response.headers["X-Cashu"] = refund_token return error_response - if ( - path.endswith("chat/completions") - or path.endswith("embeddings") - or path.endswith("messages") - or path.endswith("messages/count_tokens") - ): + if _x_cashu_path_has_settlement_handler(path): logger.debug( "Processing completion/embeddings/messages response", extra={"path": path, "amount": amount, "unit": unit}, @@ -4251,6 +4506,7 @@ class BaseUpstreamProvider: 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) @@ -4300,6 +4556,8 @@ class BaseUpstreamProvider: path: str, max_cost_for_model: int, model_obj: Model, + *, + request_body: bytes | None = None, ) -> Response | StreamingResponse: """Handle X-Cashu payment for Responses API requests. @@ -4364,6 +4622,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, mint, + request_body=request_body, ) except Exception as e: error_message = str(e) @@ -4420,6 +4679,8 @@ class BaseUpstreamProvider: max_cost_for_model: int, model_obj: Model, mint: str | None = None, + *, + request_body: bytes | None = None, ) -> Response | StreamingResponse: """Forward Responses API request paid with X-Cashu token to upstream service. @@ -4441,7 +4702,8 @@ class BaseUpstreamProvider: url = f"{self.base_url}/{path}" - request_body = await request.body() + if request_body is None: + request_body = await request.body() transformed_body = self.prepare_responses_request_body(request_body, model_obj) logger.debug( @@ -4541,6 +4803,7 @@ class BaseUpstreamProvider: 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) @@ -4592,6 +4855,7 @@ class BaseUpstreamProvider: mint: str | None = None, request_id: str | None = None, model_obj: Model | None = None, + request_body: bytes | None = None, ) -> StreamingResponse | Response: """Handle Responses API completion response for X-Cashu payment. @@ -4615,7 +4879,9 @@ class BaseUpstreamProvider: content_str = ( content.decode("utf-8") if isinstance(content, bytes) else content ) - is_streaming = content_str.startswith("data:") or "data:" in content_str + is_streaming = _is_sse_body( + response.headers.get("content-type"), content_str + ) logger.debug( "Responses API completion response analysis", @@ -4637,6 +4903,7 @@ class BaseUpstreamProvider: mint, request_id=request_id, model_obj=model_obj, + request_body=request_body, ) else: return await self.handle_x_cashu_non_streaming_responses_response( @@ -4648,6 +4915,7 @@ class BaseUpstreamProvider: mint, request_id=request_id, model_obj=model_obj, + request_body=request_body, ) except Exception as e: @@ -4676,17 +4944,20 @@ class BaseUpstreamProvider: mint: str | None = None, request_id: str | None = None, model_obj: Model | None = None, + request_body: bytes | None = None, ) -> StreamingResponse: """Handle streaming Responses API response for X-Cashu payment. Similar to regular streaming but handles Responses API specific tokens like reasoning_tokens. """ + events = _parse_sse_events(content_str) + logger.debug( "Processing streaming Responses API response", extra={ "amount": amount, "unit": unit, - "content_lines": len(content_str.strip().split("\\n")), + "event_count": len(events), }, ) @@ -4696,30 +4967,50 @@ class BaseUpstreamProvider: if "content-encoding" in response_headers: del response_headers["content-encoding"] - usage_data = None - model = None + usage_data: dict | None = None + model: str | None = None reasoning_tokens = 0 + cost_data: CostData | MaxCostData | None = None + usage_estimator = MissingUsageEstimator(request_body, model_obj) - lines = content_str.strip().split("\\n") - for line in lines: - if line.startswith("data: "): - try: - data_json = json.loads(line[6:]) - if "usage" in data_json: - usage_data = data_json["usage"] - model = data_json.get("model") - # Track reasoning tokens for Responses API - if ( - isinstance(usage_data, dict) - and "reasoning_tokens" in usage_data - ): - reasoning_tokens = usage_data.get("reasoning_tokens", 0) - elif "model" in data_json and not model: - model = data_json["model"] - except json.JSONDecodeError: - continue + for _fields, data in events: + if data.strip() == "[DONE]": + continue + try: + data_json = json.loads(data) + except json.JSONDecodeError: + continue + if not isinstance(data_json, dict): + continue + usage_estimator.observe(data_json) + # Canonical Responses API events carry model and usage nested under + # "response" (response.completed/incomplete); older shapes put them + # at the top level. + payload = _responses_usage_payload(data_json) + if isinstance(payload.get("usage"), dict): + usage_data = payload["usage"] + model = payload.get("model") or model + details = usage_data.get("output_tokens_details") + if isinstance(details, dict): + reasoning_tokens = details.get("reasoning_tokens", 0) + elif "reasoning_tokens" in usage_data: + reasoning_tokens = usage_data["reasoning_tokens"] + elif not model and payload.get("model"): + model = payload["model"] - if usage_data and model: + if not usage_data: + usage_data = usage_estimator.estimated_usage(model) + if usage_data: + logger.warning( + "No usage in streaming Responses API response, billing from local token estimate", + extra={ + "model": model, + "amount": amount, + "unit": unit, + "estimated_usage": usage_data, + }, + ) + else: logger.debug( "Found usage data in streaming Responses API response", extra={ @@ -4731,101 +5022,106 @@ class BaseUpstreamProvider: }, ) - response_data = {"usage": usage_data, "model": model} + response_data = {"usage": usage_data, "model": model or "unknown"} + try: + cost_data = await self.get_x_cashu_cost( + response_data, max_cost_for_model, model_obj + ) + if cost_data is not None and cost_data.total_msats == 0: + self._log_full_refund( + route="responses.streaming", + model=model, + content_str=content_str, + amount=amount, + unit=unit, + ) + if cost_data: + if unit == "msat": + refund_amount = amount - cost_data.total_msats + elif unit == "sat": + refund_amount = amount - (cost_data.total_msats + 999) // 1000 + else: + raise ValueError(f"Invalid unit: {unit}") + + if refund_amount > 0: + logger.debug( + "Processing refund for streaming Responses API response", + extra={ + "original_amount": amount, + "cost_msats": cost_data.total_msats, + "refund_amount": refund_amount, + "unit": unit, + "model": model, + "reasoning_tokens": reasoning_tokens, + }, + ) + + refund_token = await self.send_refund( + refund_amount, + unit, + mint, + request_id=request_id, + ) + response_headers["X-Cashu"] = refund_token + + logger.info( + "Refund processed for streaming Responses API response", + extra={ + "refund_amount": refund_amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, + }, + ) + else: + logger.debug( + "No refund needed for streaming Responses API response", + extra={ + "amount": amount, + "cost_msats": cost_data.total_msats, + "model": model, + }, + ) + + # Inject cost breakdown headers so the SDK's + # extractUsageFromResponseHeaders can populate + # inputMsats/outputMsats/totalMsats for x-cashu requests. + _inject_cost_response_headers(response_headers, cost_data) + except Exception as e: + logger.error( + "Error calculating cost for streaming Responses API response", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "model": model, + "amount": amount, + "unit": unit, + }, + ) + + for i, (fields, data) in enumerate(events): + if data.strip() == "[DONE]": + continue try: - cost_data = await self.get_x_cashu_cost( - response_data, max_cost_for_model, model_obj - ) - if cost_data: - if unit == "msat": - refund_amount = amount - cost_data.total_msats - elif unit == "sat": - refund_amount = amount - (cost_data.total_msats + 999) // 1000 - else: - raise ValueError(f"Invalid unit: {unit}") - - if refund_amount > 0: - logger.debug( - "Processing refund for streaming Responses API response", - extra={ - "original_amount": amount, - "cost_msats": cost_data.total_msats, - "refund_amount": refund_amount, - "unit": unit, - "model": model, - "reasoning_tokens": reasoning_tokens, - }, - ) - - refund_token = await self.send_refund( - refund_amount, - unit, - mint, - request_id=request_id, - ) - response_headers["X-Cashu"] = refund_token - - logger.info( - "Refund processed for streaming Responses API response", - extra={ - "refund_amount": refund_amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - else: - logger.debug( - "No refund needed for streaming Responses API response", - extra={ - "amount": amount, - "cost_msats": cost_data.total_msats, - "model": model, - }, - ) - - # Inject cost breakdown headers so the SDK's - # extractUsageFromResponseHeaders can populate - # inputMsats/outputMsats/totalMsats for x-cashu requests. - _inject_cost_response_headers(response_headers, cost_data) - except Exception as e: - logger.error( - "Error calculating cost for streaming Responses API response", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "model": model, - "amount": amount, - "unit": unit, - }, - ) - - for i, line in enumerate(lines): - if line.startswith("data: "): - try: - data_json = json.loads(line[6:]) - if not isinstance(data_json, dict): - continue - changed = False - if "provider" not in data_json: - self._apply_provider_field(data_json) - changed = True - if ( - cost_data - and "usage" in data_json - and data_json["usage"] - ): - _inject_cost_into_usage(data_json, cost_data) - changed = True - if changed: - lines[i] = "data: " + json.dumps(data_json) - except json.JSONDecodeError: - pass + data_json = json.loads(data) + except json.JSONDecodeError: + continue + if not isinstance(data_json, dict): + continue + provider_before = data_json.get("provider") + self._apply_provider_field(data_json) + changed = data_json.get("provider") != provider_before + payload = _responses_usage_payload(data_json) + if cost_data and isinstance(payload.get("usage"), dict): + _inject_cost_into_usage(payload, cost_data) + changed = True + if changed: + events[i] = (fields, json.dumps(data_json)) async def generate() -> AsyncGenerator[bytes, None]: - for line in lines: - yield (line + "\\n").encode("utf-8") + for fields, data in events: + yield _render_sse_event(fields, data).encode("utf-8") return StreamingResponse( generate(), @@ -4844,6 +5140,7 @@ class BaseUpstreamProvider: mint: str | None = None, request_id: str | None = None, model_obj: Model | None = None, + request_body: bytes | None = None, ) -> Response: """Handle non-streaming Responses API response for X-Cashu payment.""" logger.debug( @@ -4854,9 +5151,20 @@ class BaseUpstreamProvider: try: response_json = json.loads(content_str) self._apply_provider_field(response_json) + _apply_estimated_usage( + response_json, request_body, model_obj, amount, unit, "responses" + ) cost_data = await self.get_x_cashu_cost( response_json, max_cost_for_model, model_obj ) + if cost_data is not None and cost_data.total_msats == 0: + self._log_full_refund( + route="responses", + model=response_json.get("model"), + content_str=content_str, + amount=amount, + unit=unit, + ) if cost_data and "usage" in response_json: _inject_cost_into_usage(response_json, cost_data) @@ -4983,6 +5291,8 @@ class BaseUpstreamProvider: path: str, max_cost_for_model: int, model_obj: Model, + *, + request_body: bytes | None = None, ) -> Response | StreamingResponse: """Handle request with X-Cashu token payment, redeeming token and forwarding request. @@ -5007,6 +5317,21 @@ class BaseUpstreamProvider: }, ) + # Reject before redemption so the client keeps its token. + if not _x_cashu_path_has_settlement_handler(path): + logger.warning( + "Rejecting X-Cashu request for unsupported endpoint", + extra={"path": path, "method": request.method}, + ) + return create_error_response( + "invalid_request_error", + "X-Cashu payment is not supported on this endpoint; use bearer " + "(deposit) authentication instead. The token was not redeemed.", + 400, + request=request, + code="x_cashu_unsupported_endpoint", + ) + redeemed = False try: headers = dict(request.headers) @@ -5047,6 +5372,7 @@ class BaseUpstreamProvider: max_cost_for_model, model_obj, mint, + request_body=request_body, ) except Exception as e: error_message = str(e) @@ -5112,22 +5438,8 @@ class BaseUpstreamProvider: {k: v * self.provider_fee for k, v in base_pricing.dict().items()} ) - temp_model = Model( - id=model.id, - name=model.name, - created=model.created, - description=model.description, - context_length=model.context_length, - architecture=model.architecture, - pricing=adjusted_pricing, - sats_pricing=None, - per_request_limits=model.per_request_limits, - top_provider=model.top_provider, - enabled=model.enabled, - upstream_provider_id=model.upstream_provider_id, - canonical_slug=model.canonical_slug, - alias_ids=model.alias_ids, - forwarded_model_id=model.forwarded_model_id, + temp_model = model.copy( + update={"pricing": adjusted_pricing, "sats_pricing": None} ) ( @@ -5136,23 +5448,7 @@ class BaseUpstreamProvider: adjusted_pricing.max_cost, ) = _calculate_usd_max_costs(temp_model) - return Model( - id=model.id, - name=model.name, - created=model.created, - description=model.description, - context_length=model.context_length, - architecture=model.architecture, - pricing=adjusted_pricing, - sats_pricing=model.sats_pricing, - per_request_limits=model.per_request_limits, - top_provider=model.top_provider, - enabled=model.enabled, - upstream_provider_id=model.upstream_provider_id, - canonical_slug=model.canonical_slug, - alias_ids=model.alias_ids, - forwarded_model_id=model.forwarded_model_id, - ) + return model.copy(update={"pricing": adjusted_pricing}) async def fetch_models(self) -> list[Model]: """Fetch available models from upstream API and update cache. @@ -5311,7 +5607,7 @@ class BaseUpstreamProvider: except Exception as e: logger.error( f"Failed to refresh models cache for {self.provider_type or self.base_url}", - extra={"error": str(e), "error_type": type(e).__name__}, + extra={"error": repr(e), "error_type": type(e).__name__}, ) def get_cached_models(self) -> list[Model]: diff --git a/routstr/upstream/count_tokens.py b/routstr/upstream/count_tokens.py index d114561c..ea067dd1 100644 --- a/routstr/upstream/count_tokens.py +++ b/routstr/upstream/count_tokens.py @@ -23,7 +23,11 @@ import litellm from fastapi.responses import Response from ..core import get_logger -from ..payment.helpers import estimate_tokens +from ..payment.helpers import ( + _count_prompt_token_ids, + estimate_prompt_tokens, + estimate_tokens, +) from ..payment.models import Model logger = get_logger(__name__) @@ -39,11 +43,66 @@ def _parse_request_body(request_body: bytes | None) -> dict[str, Any]: return parsed if isinstance(parsed, dict) else {} -def _count_with_litellm(model: str, body: dict[str, Any]) -> int: +def _model_name(model_obj: Model | None, body: dict[str, Any]) -> str: + if model_obj is not None: + return model_obj.forwarded_model_id or model_obj.id or "" + body_model = body.get("model") + return body_model if isinstance(body_model, str) else "" + + +def _count_with_litellm( + model: str, body: dict[str, Any], include_legacy_prompt: bool = False +) -> int: messages = body.get("messages") if not isinstance(messages, list): messages = [] + if "input" in body: + response_input = body["input"] + if isinstance(response_input, str): + messages = [{"role": "user", "content": response_input}] + elif isinstance(response_input, list): + messages = [] + for item in response_input: + if not isinstance(item, dict) or "role" not in item: + raise ValueError( + "Responses input requires fallback token estimation" + ) + content = item.get("content", "") + if isinstance(content, list): + parts = [] + for part in content: + if not isinstance(part, dict) or part.get("type") not in ( + "input_text", + "output_text", + "text", + ): + raise ValueError( + "Non-text Responses input requires fallback token estimation" + ) + parts.append({"type": "text", "text": part.get("text", "")}) + content = parts + messages.append({"role": item["role"], "content": content}) + else: + raise ValueError("Unsupported Responses input") + if body.get("instructions"): + messages.insert(0, {"role": "system", "content": body["instructions"]}) + + prompt_token_ids = 0 + if include_legacy_prompt: + prompt = body.get("prompt") + if isinstance(prompt, str): + prompt_texts = [prompt] + elif isinstance(prompt, list): + prompt_texts = [item for item in prompt if isinstance(item, str)] + else: + prompt_texts = [] + prompt_token_ids = _count_prompt_token_ids(prompt) + messages = [ + *({"role": "user", "content": text} for text in prompt_texts if text), + *messages, + ] + system = body.get("system") if isinstance(system, str) and system: messages = [{"role": "system", "content": system}, *messages] @@ -58,7 +117,7 @@ def _count_with_litellm(model: str, body: dict[str, Any]) -> int: tools = body.get("tools") if isinstance(body.get("tools"), list) else None - return int( + return prompt_token_ids + int( litellm.token_counter( model=model, messages=messages, @@ -67,6 +126,173 @@ def _count_with_litellm(model: str, body: dict[str, Any]) -> int: ) +def _count_text_with_litellm(model: str, text: str) -> int: + return int( + litellm.token_counter( + model=model, + text=text, + count_response_tokens=True, + ) + ) + + +def _generated_text(value: object) -> list[str]: + """Extract generated text/tool arguments without counting response metadata.""" + generated_keys = { + "arguments", + "content", + "delta", + "output_text", + "partial_json", + "reasoning", + "reasoning_content", + "text", + "thinking", + } + parts: list[str] = [] + + def walk(item: object, key: str | None = None) -> None: + if isinstance(item, str): + if key in generated_keys: + parts.append(item) + return + if isinstance(item, list): + for child in item: + walk(child, key) + return + if isinstance(item, dict): + for child_key, child in item.items(): + walk(child, child_key) + + walk(value) + return parts + + +class MissingUsageEstimator: + """Estimate billable usage when an upstream omits its usage trailer. + + The reservation is deliberately absent from this class: it is an + authorization ceiling, not an input to usage measurement. + """ + + def __init__(self, request_body: bytes | None, model_obj: Model | None) -> None: + self.body = _parse_request_body(request_body) + self.model_name = _model_name(model_obj, self.body) + self._output_parts: list[str] = [] + self._input_tokens: int | None = None + + def _estimate_input_tokens(self) -> int: + if self._input_tokens is not None: + return self._input_tokens + try: + self._input_tokens = _count_with_litellm( + self.model_name, self.body, include_legacy_prompt=True + ) + except Exception as exc: + self._input_tokens = estimate_prompt_tokens(self.body) + logger.debug( + "litellm request token count failed; using local estimator", + extra={ + "model": self.model_name, + "error": str(exc), + "error_type": type(exc).__name__, + "estimated_tokens": self._input_tokens, + }, + ) + return self._input_tokens + + @property + def output_text(self) -> str: + return "".join(self._output_parts) + + def observe(self, response_data: object) -> None: + if isinstance(response_data, dict): + event_type = response_data.get("type") + if event_type in ("response.completed", "response.incomplete"): + response = response_data.get("response") + if isinstance(response, dict) and isinstance( + response.get("output"), list + ): + # Terminal output is a snapshot, not another text delta. + self._output_parts = _generated_text(response["output"]) + return + if isinstance(event_type, str) and event_type.endswith(".done"): + # Responses API ``*.done`` events repeat text already streamed + # via ``*.delta`` events; counting both would double-bill. + return + self._output_parts.extend(_generated_text(response_data)) + + def estimated_usage(self, model: str | None = None) -> dict[str, Any] | None: + """Local usage estimate, or None when the upstream generated no text.""" + if not self.output_text: + return None + return self.response_data(model)["usage"] + + def billing_data( + self, + response_data: dict[str, Any] | None, + model: str | None = None, + ) -> dict[str, Any]: + """Use measured usage when present, otherwise return a local estimate.""" + if isinstance(response_data, dict): + usage = response_data.get("usage") + if not isinstance(usage, dict): + nested = response_data.get("response") + usage = nested.get("usage") if isinstance(nested, dict) else None + if isinstance(usage, dict) and usage: + return { + "model": model or response_data.get("model") or self.model_name, + "usage": usage, + } + if not self._output_parts: + self.observe(response_data) + return self.response_data(model) + + def response_data(self, model: str | None = None) -> dict[str, Any]: + text = self.output_text + try: + output_tokens = ( + _count_text_with_litellm(self.model_name, text) if text else 0 + ) + except Exception as exc: + output_tokens = len(text) // 3 + logger.debug( + "litellm response token count failed; using local estimator", + extra={ + "model": self.model_name, + "error": str(exc), + "error_type": type(exc).__name__, + "estimated_tokens": output_tokens, + }, + ) + + input_tokens = max(0, int(self._estimate_input_tokens())) + output_tokens = max(0, int(output_tokens)) + return { + "model": model or self.model_name or "unknown", + "usage": { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + "total_tokens": input_tokens + output_tokens, + "estimated": True, + }, + } + + def openai_response_data(self, model: str | None = None) -> dict[str, Any]: + """Same estimate in the OpenAI chat-completions usage dialect.""" + data = self.response_data(model) + usage = data["usage"] + return { + "model": data["model"], + "usage": { + "prompt_tokens": usage["input_tokens"], + "completion_tokens": usage["output_tokens"], + "total_tokens": usage["total_tokens"], + "estimated": True, + }, + } + + def count_tokens_locally( request_body: bytes | None, model_obj: Model | None, @@ -75,13 +301,7 @@ def count_tokens_locally( touching the upstream. Always returns 200; never raises.""" body = _parse_request_body(request_body) - model_name = "" - if model_obj is not None: - model_name = model_obj.forwarded_model_id or model_obj.id or "" - if not model_name: - body_model = body.get("model") - if isinstance(body_model, str): - model_name = body_model + model_name = _model_name(model_obj, body) input_tokens: int try: diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index c70a3b68..90b92719 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -10,17 +10,17 @@ from urllib.parse import urlsplit, urlunsplit from fastapi import Request from fastapi.responses import Response, StreamingResponse -from sqlalchemy import case -from sqlmodel import col, update from ..auth import ( ROUTSTR_FEE_PERCENT, ReservationSnapshot, + _charge_reservation_rows, _claim_reservation_for_charge, + _stop_reservation_heartbeat, _validate_reservation_snapshot, - get_billing_key, get_reservation_snapshot, payments_logger, + release_reservation, ) from ..core import get_logger from ..core.db import ( @@ -31,7 +31,7 @@ from ..core.db import ( from ..core.db import ( store_cashu_transaction_with_retry as store_cashu_transaction, ) -from ..core.exceptions import UpstreamError +from ..core.exceptions import EhbpTimeoutError, UpstreamError from ..core.settings import settings from ..payment.cost_calculation import ( CostData, @@ -40,8 +40,13 @@ from ..payment.cost_calculation import ( ) from ..payment.helpers import create_error_response from ..payment.models import Model -from ..wallet import recieve_token, send_token -from .tinfoil_trailer import forward_with_trailer +from ..wallet import ( + SPENT_TOKEN_CODES, + classify_redemption_error, + recieve_token, + send_token, +) +from .tinfoil_trailer import TrailerResponse, forward_with_trailer logger = get_logger(__name__) @@ -57,6 +62,55 @@ _TINFOIL_ALLOWED_ENCLAVE_HOST_SUFFIX = ".tinfoil.sh" _TINFOIL_ALLOWED_ENCLAVE_HOSTS = frozenset({"tinfoil.sh"}) +_KEY_CONFIG_PROBLEM_TYPE = "urn:ietf:params:ehbp:error:key-config" + + +def _is_ehbp_key_config_response(resp: TrailerResponse) -> bool: + """Check whether an upstream EHBP response is a key-config mismatch. + + The enclave returns ``422 application/problem+json`` with + ``type=urn:ietf:params:ehbp:error:key-config`` when it cannot decrypt the + request body — meaning the client's HPKE key is stale (the enclave rotated + keys). The proxy must pass this response through with its original + content type so EHBP clients can detect it and trigger re-attestation. + """ + if resp.status_code != 422: + return False + ct = "" + for k, v in resp.headers: + if k.lower() == "content-type": + ct = v.lower() + break + media_type = ct.split(";", 1)[0].strip() + if media_type != "application/problem+json": + return False + try: + body = json.loads(resp.body) + except (json.JSONDecodeError, UnicodeDecodeError): + return False + return isinstance(body, dict) and body.get("type") == _KEY_CONFIG_PROBLEM_TYPE + + +def _passthrough_key_config_response(resp: TrailerResponse) -> Response: + """Return the enclave's key-config 422 with its original body and content + type so the EHBP client's ``KeyConfigMismatchError`` detection fires. + + Only the content type is forwarded. ``Ehbp-Response-Nonce`` must be + dropped: a nonce only carries meaning for an *encrypted* response body, + and the stock ``ehbp`` client (``shouldDecryptResponse``) checks for the + nonce *before* checking for a key-config mismatch — forwarding it would + send that client down the decrypt path on this plaintext error body, so + the re-attestation loop would never fire. Content-length is recomputed + from the body, and upstream-internal headers are filtered out. + """ + return Response( + content=resp.body, + status_code=422, + headers={"content-type": "application/problem+json"}, + media_type="application/problem+json", + ) + + def _normalize_upstream_model_id(model_id: str | None) -> str: """Normalize casing and whitespace for upstream identity comparisons.""" if not model_id: @@ -73,28 +127,43 @@ _PROXY_ONLY_HEADERS = frozenset( } ) +# Namespace prefix the routstr catalog applies to Tinfoil models +# (e.g. ``tinfoil-deepseek-v4-1-flash``). The SDK strips this prefix for the +# encrypted body (``getTinfoilUpstreamModelId`` in client/TinfoilSecure.ts), so +# the enclave always reports the *bare* upstream model id in the usage-metrics +# header even though the routstr model id and ``forwarded_model_id`` carry it. +TINFOIL_MODEL_PREFIX = "tinfoil-" + def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: """Parse ``X-Tinfoil-Usage-Metrics`` into an OpenAI-style usage dict. The header format is:: - prompt=,completion=,total=[,model=] + prompt=,completion=,total=[,cached_prompt_tokens=, + uncached_prompt_tokens=][,model=][,cost_usd=] + + ``prompt`` is the inclusive prompt total and ``cached_prompt_tokens`` is + the cache-read portion included within it. Routstr maps these to + ``prompt_tokens`` and ``cache_read_input_tokens`` so ``normalize_usage`` + can subtract the cached read from the prompt total (OpenAI-family + semantics). ``cost_usd`` is parsed as a float and kept for logging/ + cross-checking only — billing uses the token path. The ``model`` field (added in tinfoilsh/confidential-model-router PR #385) - is extracted as a string and included in the returned dict under the - ``"model"`` key so callers can compare the served model against the - requested one and adjust pricing. + is extracted as a string so callers can compare the served model against + the requested one and adjust pricing. - Returns a dict like ``{"prompt_tokens": n, "completion_tokens": n, - "model": ""}`` suitable for :func:`calculate_cost` (which ignores - the extra ``model`` key in the usage sub-dict), or ``None`` when the + Returns a dict suitable for :func:`calculate_cost`, or ``None`` when the header is absent or malformed. """ if not header_value: return None - parts: dict[str, int] = {} + + int_parts: dict[str, int] = {} model: str | None = None + cost_usd: float | None = None + for item in header_value.split(","): key, sep, value = item.partition("=") if not sep: @@ -104,30 +173,44 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None: if key == "model": model = value continue + if key == "cost_usd": + try: + cost_usd = float(value) + except (ValueError, TypeError): + cost_usd = None + continue try: - parts[key] = int(value) + int_parts[key] = int(value) except (ValueError, TypeError): continue - prompt = parts.get("prompt") - completion = parts.get("completion") - if prompt is not None and completion is not None: - result: dict[str, int | str] = { - "prompt_tokens": prompt, - "completion_tokens": completion, - } - if "total" in parts: - result["total_tokens"] = parts["total"] - if model: - result["model"] = model - return result - logger.warning( - "Failed to parse X-Tinfoil-Usage-Metrics header", - extra={ - "header_value": header_value, - "parsed_parts": parts, - }, - ) - return None + + prompt = int_parts.get("prompt") + completion = int_parts.get("completion") + if prompt is None or completion is None: + logger.warning( + "Failed to parse X-Tinfoil-Usage-Metrics header", + extra={ + "header_value": header_value, + "parsed_parts": int_parts, + }, + ) + return None + + result: dict[str, int | float | str] = { + "prompt_tokens": prompt, + "completion_tokens": completion, + } + if "total" in int_parts: + result["total_tokens"] = int_parts["total"] + if "cached_prompt_tokens" in int_parts: + result["cache_read_input_tokens"] = int_parts["cached_prompt_tokens"] + if "uncached_prompt_tokens" in int_parts: + result["uncached_prompt_tokens"] = int_parts["uncached_prompt_tokens"] + if cost_usd is not None: + result["cost_usd"] = cost_usd + if model: + result["model"] = model + return result def _get_header_case_insensitive( @@ -191,7 +274,9 @@ def _resolve_ehbp_target_url( otherwise the header is ignored so callers cannot redirect other providers or leak upstream API keys. """ - override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER + override_header = ( + profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER + ) if not override_header: return target_url enclave_url = _get_header_case_insensitive(headers, override_header) @@ -274,6 +359,11 @@ def _build_cost_info( output_tokens: int = 0, input_msats: int = 0, output_msats: int = 0, + cache_read_input_tokens: int = 0, + cache_creation_input_tokens: int = 0, + cache_read_msats: int = 0, + cache_creation_msats: int = 0, + total_usd: float = 0.0, actual_model: str | None = None, ) -> dict: """Build a cost-info dict with token counts and per-token-type costs. @@ -282,22 +372,25 @@ def _build_cost_info( one), it is included in the returned dict so callers can use it for billing finalization and logging. """ - result: dict[str, int | str | None] = { + result: dict[str, int | float | str | None] = { "total_msats": total_msats, "input_tokens": input_tokens, "output_tokens": output_tokens, "total_tokens": input_tokens + output_tokens, "input_msats": input_msats, "output_msats": output_msats, + "cache_read_input_tokens": cache_read_input_tokens, + "cache_creation_input_tokens": cache_creation_input_tokens, + "cache_read_msats": cache_read_msats, + "cache_creation_msats": cache_creation_msats, + "total_usd": total_usd, } if actual_model: result["actual_model"] = actual_model return result -def _inject_cost_response_headers( - headers: dict[str, str], cost_info: dict -) -> None: +def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None: """Add per-request cost headers to an EHBP response. Since EHBP response bodies are opaque encrypted blobs, cost cannot be @@ -305,8 +398,14 @@ def _inject_cost_response_headers( the client/Tinfoil SDK can read without decrypting. """ headers["X-Routstr-Cost-Msats"] = str(cost_info["total_msats"]) + if "computed_msats" in cost_info: + headers["X-Routstr-Computed-Cost-Msats"] = str(cost_info["computed_msats"]) headers["X-Routstr-Input-Cost-Msats"] = str(cost_info["input_msats"]) headers["X-Routstr-Output-Cost-Msats"] = str(cost_info["output_msats"]) + headers["X-Routstr-Cache-Read-Msats"] = str(cost_info.get("cache_read_msats", 0)) + headers["X-Routstr-Cache-Creation-Msats"] = str( + cost_info.get("cache_creation_msats", 0) + ) async def _compute_ehbp_actual_cost( @@ -316,10 +415,10 @@ async def _compute_ehbp_actual_cost( ) -> dict: """Compute the actual cost in msats from Tinfoil usage metrics. - Falls back to ``max_cost_for_model`` when usage is absent (streaming) or - cannot be priced. The result is clamped to ``[min_request_msat, - max_cost_for_model]`` so the refund never exceeds the reservation and is - never zero. + When usage is present, the result is clamped to ``[min_request_msat, + max_cost_for_model]``. Missing or unpriceable usage returns zero: encrypted + EHBP bodies cannot be estimated locally, and the authorization ceiling is + not evidence of consumption. When the usage-metrics header includes ``model=`` and it differs from ``model_obj.id``, the actual served model's pricing is used for the @@ -332,7 +431,7 @@ async def _compute_ehbp_actual_cost( """ usage_dict = parse_tinfoil_usage_metrics(usage_header) if usage_dict is None: - return _build_cost_info(max_cost_for_model) + return _build_cost_info(0) # The enclave may serve a different model than the one requested (e.g. # due to failover). The usage-metrics header's ``model=`` carries @@ -344,6 +443,14 @@ async def _compute_ehbp_actual_cost( # look up the actual model's pricing. actual_model: str | None = usage_dict.pop("model", None) # type: ignore[arg-type] pricing_model_id = model_obj.id + # Bill the model we actually routed to. Passing only the model *string* + # to calculate_cost makes it re-derive pricing from the global alias map, + # which resolves the id to the best-ranked candidate — not the serving + # one. Tinfoil's catalog id (e.g. ``deepseek-v4-1-flash``) is also a + # cross-provider alias, and that cheaper candidate has no cache rate, so + # the cache discount silently disappeared (and the request was + # undercharged). Hand calculate_cost the identity it cannot reconstruct. + pricing_model_obj: Model = model_obj expected_upstream_model = model_obj.forwarded_model_id or model_obj.id expected_identity = _normalize_upstream_model_id(expected_upstream_model) served_identity = _normalize_upstream_model_id(actual_model) @@ -359,7 +466,24 @@ async def _compute_ehbp_actual_cost( # the global model map. The resolved object can belong to a different # provider and therefore have a different client-facing ``id`` while # still representing the same upstream model. - actual_model_obj = get_model_instance(actual_model) + # + # The enclave reports the *bare* upstream id, but the routstr model is + # namespaced ``tinfoil-`` (and the SDK strips that prefix for the + # encrypted body). Resolve the served id within the same namespace + # first: a same-model report then maps back onto the requested Tinfoil + # model, and a genuine failover lands on the actually-served Tinfoil + # model — instead of the cheaper cross-provider model the bare id + # would resolve to in the global map. + namespaced_served = actual_model + if ( + expected_upstream_model.startswith(TINFOIL_MODEL_PREFIX) + and not actual_model.startswith(TINFOIL_MODEL_PREFIX) + ): + namespaced_served = TINFOIL_MODEL_PREFIX + actual_model + + actual_model_obj = get_model_instance(namespaced_served) + if actual_model_obj is None and namespaced_served != actual_model: + actual_model_obj = get_model_instance(actual_model) if actual_model_obj is None: logger.warning( "EHBP served model not found in registry, falling back " @@ -375,9 +499,7 @@ async def _compute_ehbp_actual_cost( resolved_upstream_model = ( actual_model_obj.forwarded_model_id or actual_model_obj.id ) - resolved_identity = _normalize_upstream_model_id( - resolved_upstream_model - ) + resolved_identity = _normalize_upstream_model_id(resolved_upstream_model) if resolved_identity != expected_identity: logger.info( "EHBP served model differs from requested, using actual " @@ -390,6 +512,7 @@ async def _compute_ehbp_actual_cost( }, ) pricing_model_id = actual_model_obj.id + pricing_model_obj = actual_model_obj else: # A different registry/client alias resolved to the same # upstream model; retain the requested model's pricing. @@ -402,22 +525,23 @@ async def _compute_ehbp_actual_cost( cost = await calculate_cost( {"model": pricing_model_id, "usage": usage_dict}, max_cost_for_model, + pricing_model_obj, ) except Exception as e: logger.warning( - "EHBP usage cost calculation failed, falling back to max cost", + "EHBP usage cost calculation failed; releasing instead of charging max cost", extra={ "model": pricing_model_id, "error": str(e), "usage": usage_dict, }, ) - return _build_cost_info(max_cost_for_model, actual_model=actual_model) + return _build_cost_info(0, actual_model=actual_model) if isinstance(cost, MaxCostData): logger.warning( - "EHBP calculate_cost returned MaxCostData (no model pricing), " - "falling back to max cost", + "EHBP calculate_cost returned MaxCostData (no usable pricing); " + "releasing instead of charging max cost", extra={ "model": pricing_model_id, "max_cost_for_model": max_cost_for_model, @@ -425,7 +549,7 @@ async def _compute_ehbp_actual_cost( "cost_total_msats": cost.total_msats, }, ) - return _build_cost_info(max_cost_for_model, actual_model=actual_model) + return _build_cost_info(0, actual_model=actual_model) if isinstance(cost, CostData): actual = max(int(cost.total_msats), int(settings.min_request_msat)) clamped = min(actual, max_cost_for_model) @@ -445,17 +569,22 @@ async def _compute_ehbp_actual_cost( output_tokens=cost.output_tokens, input_msats=cost.input_msats, output_msats=cost.output_msats, + cache_read_input_tokens=cost.cache_read_input_tokens, + cache_creation_input_tokens=cost.cache_creation_input_tokens, + cache_read_msats=cost.cache_read_msats, + cache_creation_msats=cost.cache_creation_msats, + total_usd=cost.total_usd, actual_model=actual_model, ) # CostDataError logger.warning( - "EHBP usage cost calculation error, falling back to max cost", + "EHBP usage cost calculation error; releasing instead of charging max cost", extra={ "model": pricing_model_id, "error": getattr(cost, "message", str(cost)), }, ) - return _build_cost_info(max_cost_for_model, actual_model=actual_model) + return _build_cost_info(0, actual_model=actual_model) def _extract_usage_from_response( @@ -500,6 +629,18 @@ class EHBPForwardingTarget: profile: ConfidentialInferenceProfile | None = None +async def _release_failed_ehbp_charge( + reservation: ReservationSnapshot, session: AsyncSession +) -> None: + if await release_reservation(reservation, session, reservation.reserved_msats): + return + await _stop_reservation_heartbeat(reservation.release_id) + logger.critical( + "Failed to release EHBP reservation after rejected charge", + extra={"reservation_id": reservation.release_id}, + ) + + async def finalize_ehbp_actual_cost_payment( key: ApiKey, session: AsyncSession, @@ -507,61 +648,27 @@ async def finalize_ehbp_actual_cost_payment( model_id: str, cost_info: dict, reservation_snapshot: ReservationSnapshot | None = None, -) -> None: +) -> int: """Finalize an EHBP bearer request using clamped provider usage metrics.""" reservation = reservation_snapshot or await get_reservation_snapshot(key, session) await _validate_reservation_snapshot(key, reservation, session) if not await _claim_reservation_for_charge(reservation, session): - return + return 0 reserved_cost_for_model = reservation.reserved_msats - billing_key = await get_billing_key(key, session) key_hash = key.hashed_key - billing_key_hash = billing_key.hashed_key - total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model))) + billing_key_hash = key_hash + total_cost_msats = max( + 0, int(cost_info.get("total_msats", reserved_cost_for_model)) + ) now = int(time.time()) - safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= reserved_cost_for_model, - col(ApiKey.reserved_balance) - reserved_cost_for_model, - ), - else_=0, + charged = await _charge_reservation_rows( + session, + billing_key_hash=billing_key_hash, + reserved_msats=reserved_cost_for_model, + charge_msats=total_cost_msats, ) - cleared_reserved_at = case( - ( - col(ApiKey.reserved_balance) - reserved_cost_for_model > 0, - col(ApiKey.reserved_at), - ), - else_=None, - ) - - stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == billing_key.hashed_key) - .values( - reserved_balance=safe_reserved, - reserved_at=cleared_reserved_at, - balance=col(ApiKey.balance) - total_cost_msats, - total_spent=col(ApiKey.total_spent) + total_cost_msats, - ) - ) - result = await session.exec(stmt) # type: ignore[call-overload] - - child_result = None - if billing_key.hashed_key != key.hashed_key: - child_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values( - reserved_balance=safe_reserved, - reserved_at=cleared_reserved_at, - total_spent=col(ApiKey.total_spent) + total_cost_msats, - ) - ) - child_result = await session.exec(child_stmt) # type: ignore[call-overload] - - if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): - await session.rollback() + if not charged: logger.error( "Failed to finalize EHBP usage-based payment", extra={ @@ -570,16 +677,14 @@ async def finalize_ehbp_actual_cost_payment( "model": model_id, "reserved_cost_for_model": reserved_cost_for_model, "total_cost_msats": total_cost_msats, - "parent_rowcount": result.rowcount, - "child_rowcount": getattr(child_result, "rowcount", None), }, ) - return + await _release_failed_ehbp_charge(reservation, session) + return 0 await session.commit() - await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) + await _stop_reservation_heartbeat(reservation.release_id) + await session.refresh(key) if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0: fee_msats = math.ceil(total_cost_msats * ROUTSTR_FEE_PERCENT / 100) @@ -596,19 +701,30 @@ async def finalize_ehbp_actual_cost_payment( extra={ "event": "finalize", "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", + "billing_key_hash": key.hashed_key[:8] + "...", "model": model_id, "cost_reserved": reserved_cost_for_model, "cost_charged": total_cost_msats, "input_tokens": cost_info.get("input_tokens", 0), "output_tokens": cost_info.get("output_tokens", 0), - "balance": billing_key.balance, - "reserved_balance": billing_key.reserved_balance, - "total_spent": billing_key.total_spent, + # Cache splits are only knowable when the enclave reports + # ``cached_prompt_tokens``; absent that they are a measured zero on + # the token counts the provider did report (not an unknown), so the + # event key set stays stable for usage-analytics consumers. + "cache_read_input_tokens": cost_info.get("cache_read_input_tokens", 0), + "cache_creation_input_tokens": cost_info.get( + "cache_creation_input_tokens", 0 + ), + "cache_read_msats": cost_info.get("cache_read_msats", 0), + "cache_creation_msats": cost_info.get("cache_creation_msats", 0), + "balance": key.balance, + "reserved_balance": key.reserved_balance, + "total_spent": key.total_spent, "finalize_type": "ehbp_usage", "finalized_at": now, }, ) + return total_cost_msats async def finalize_ehbp_max_cost_payment( @@ -617,128 +733,26 @@ async def finalize_ehbp_max_cost_payment( max_cost_for_model: int, model_id: str, reservation_snapshot: ReservationSnapshot | None = None, -) -> None: - """Finalize an EHBP bearer request by charging the reserved max cost. +) -> int: + """Release an unmeasured EHBP request without charging its reservation. - EHBP responses are encrypted, so Routstr cannot inspect token usage. Unlike - normal completion handlers, this intentionally charges the pre-reserved max - cost and releases the reservation. + The legacy name is retained for compatibility with internal callers. EHBP + responses are encrypted, so no local estimate is possible when the trusted + usage header/trailer is absent. """ reservation = reservation_snapshot or await get_reservation_snapshot(key, session) await _validate_reservation_snapshot(key, reservation, session) - if not await _claim_reservation_for_charge(reservation, session): - return - max_cost_for_model = reservation.reserved_msats - billing_key = await get_billing_key(key, session) - key_hash = key.hashed_key - billing_key_hash = billing_key.hashed_key - total_cost_msats = max(0, int(max_cost_for_model)) - now = int(time.time()) - - cleared_reserved_at = case( - ( - col(ApiKey.reserved_balance) - max_cost_for_model > 0, - col(ApiKey.reserved_at), - ), - else_=None, - ) - safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= max_cost_for_model, - col(ApiKey.reserved_balance) - max_cost_for_model, - ), - else_=0, - ) - - stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == billing_key.hashed_key) - .values( - reserved_balance=safe_reserved, - reserved_at=cleared_reserved_at, - balance=col(ApiKey.balance) - total_cost_msats, - total_spent=col(ApiKey.total_spent) + total_cost_msats, - ) - ) - result = await session.exec(stmt) # type: ignore[call-overload] - - if billing_key.hashed_key != key.hashed_key: - child_safe_reserved = case( - ( - col(ApiKey.reserved_balance) >= max_cost_for_model, - col(ApiKey.reserved_balance) - max_cost_for_model, - ), - else_=0, - ) - child_cleared_reserved_at = case( - ( - col(ApiKey.reserved_balance) - max_cost_for_model > 0, - col(ApiKey.reserved_at), - ), - else_=None, - ) - child_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .values( - reserved_balance=child_safe_reserved, - reserved_at=child_cleared_reserved_at, - total_spent=col(ApiKey.total_spent) + total_cost_msats, - ) - ) - child_result = await session.exec(child_stmt) # type: ignore[call-overload] - else: - child_result = None - - if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): - await session.rollback() - logger.error( - "Failed to finalize EHBP max-cost payment", - extra={ - "key_hash": key_hash[:8] + "...", - "billing_key_hash": billing_key_hash[:8] + "...", - "model": model_id, - "max_cost_for_model": max_cost_for_model, - "parent_rowcount": result.rowcount, - "child_rowcount": getattr(child_result, "rowcount", None), - }, - ) - return - - await session.commit() - - await session.refresh(billing_key) - if billing_key.hashed_key != key.hashed_key: - await session.refresh(key) - - if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0: - fee_msats = math.ceil(total_cost_msats * ROUTSTR_FEE_PERCENT / 100) - try: - await accumulate_routstr_fee(session, fee_msats) - except Exception as e: - logger.warning( - "Failed to accumulate Routstr fee for EHBP request", - extra={"error": str(e), "fee_msats": fee_msats}, - ) - - payments_logger.info( - "FINALIZE", + key_log_hash = key.hashed_key[:8] + "..." + await release_reservation(reservation, session, reservation.reserved_msats) + logger.warning( + "Released unmeasured EHBP reservation without charging max cost", extra={ - "event": "finalize", - "key_hash": key.hashed_key[:8] + "...", - "billing_key_hash": billing_key.hashed_key[:8] + "...", + "key_hash": key_log_hash, "model": model_id, - "cost_reserved": max_cost_for_model, - "cost_charged": total_cost_msats, - "input_tokens": 0, - "output_tokens": 0, - "balance": billing_key.balance, - "reserved_balance": billing_key.reserved_balance, - "total_spent": billing_key.total_spent, - "finalize_type": "ehbp_max_cost", - "finalized_at": now, + "max_cost_for_model": max_cost_for_model, }, ) + return 0 async def send_cashu_refund( @@ -842,6 +856,25 @@ async def forward_ehbp_request( "body_preview": body_preview, }, ) + # Key-config mismatch (stale client HPKE key): return the + # enclave's 422 problem+json directly so the SDK's + # KeyConfigMismatchError detection fires and triggers + # re-attestation. Wrapping it as application/json would + # destroy the signal and cause permanent failure. + if _is_ehbp_key_config_response(resp): + logger.warning( + "EHBP upstream %s returned key-config mismatch for model=%s, " + "passing through for client re-attestation", + provider_type, + model_obj.id, + extra={ + "provider": provider_type, + "model": model_obj.id, + "path": path, + }, + ) + return _passthrough_key_config_response(resp) + raise UpstreamError( f"EHBP upstream {provider_type} returned {resp.status_code} " f"for model {model_obj.id}: {body_preview[:200] or ''}", @@ -893,10 +926,9 @@ async def forward_ehbp_request( cost_info = await _compute_ehbp_actual_cost( usage_header, model_obj, max_cost_for_model ) - # Use the actual served model for billing when it differs from - # the requested model. billing_model = cost_info.pop("actual_model", None) or model_obj.id - await finalize_ehbp_actual_cost_payment( + computed_msats = int(cost_info["total_msats"]) + charged_msats = await finalize_ehbp_actual_cost_payment( key, session, max_cost_for_model, @@ -904,18 +936,25 @@ async def forward_ehbp_request( cost_info, reservation_snapshot, ) - cost_data = {**cost_info, "total_usd": 0.0} + cost_data = { + **cost_info, + "total_msats": charged_msats, + "charged_msats": charged_msats, + "total_usd": cost_info.get("total_usd", 0.0), + } + if computed_msats != charged_msats: + cost_data["computed_msats"] = computed_msats else: logger.warning( - "EHBP usage metrics not found in headers or trailers, " - "falling back to max-cost billing", + "EHBP usage metrics not found in headers or trailers; " + "releasing instead of charging the authorization ceiling", extra={ "model": model_obj.id, "provider": provider_type, "key_hash": key.hashed_key[:8] + "...", }, ) - await finalize_ehbp_max_cost_payment( + charged_msats = await finalize_ehbp_max_cost_payment( key, session, max_cost_for_model, @@ -923,14 +962,15 @@ async def forward_ehbp_request( reservation_snapshot, ) cost_data = { - "total_msats": max_cost_for_model, + "total_msats": charged_msats, + "charged_msats": charged_msats, "total_usd": 0.0, "input_tokens": 0, "output_tokens": 0, } - # Build the cost_info dict from what adjust_payment_for_tokens returned - # or from the max-cost fallback. Fields match CostData/MaxCostData.dict(). + # Build the cost_info dict from measured usage or the unmeasured-release + # fallback. Fields match CostData/MaxCostData.dict(). cost_info = { "total_msats": cost_data.get("total_msats", max_cost_for_model), "input_tokens": cost_data.get("input_tokens", 0), @@ -939,7 +979,15 @@ async def forward_ehbp_request( + cost_data.get("output_tokens", 0), "input_msats": cost_data.get("input_msats", 0), "output_msats": cost_data.get("output_msats", 0), + "cache_read_input_tokens": cost_data.get("cache_read_input_tokens", 0), + "cache_creation_input_tokens": cost_data.get( + "cache_creation_input_tokens", 0 + ), + "cache_read_msats": cost_data.get("cache_read_msats", 0), + "cache_creation_msats": cost_data.get("cache_creation_msats", 0), } + if "computed_msats" in cost_data: + cost_info["computed_msats"] = cost_data["computed_msats"] cost_usd = cost_data.get("total_usd", 0.0) # Build response headers, filtering out hop-by-hop headers @@ -1028,7 +1076,9 @@ async def forward_ehbp_x_cashu_request( target_url = _resolve_ehbp_target_url( target.url, path, headers, provider_type, profile ) - upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile) + upstream_headers = _prepare_ehbp_upstream_headers( + headers, target.headers, profile + ) request_body = await request.body() # Merge query params into the target URL @@ -1047,6 +1097,29 @@ async def forward_ehbp_x_cashu_request( ) if resp.status_code != 200: + # Key-config mismatch (stale client HPKE key): refund the + # full token and pass the enclave's 422 problem+json through + # so the SDK's KeyConfigMismatchError detection fires. + if _is_ehbp_key_config_response(resp): + logger.warning( + "EHBP upstream %s returned key-config mismatch for " + "model=%s, refunding and passing through", + provider_type, + model_obj.id, + extra={ + "provider": provider_type, + "model": model_obj.id, + "path": path, + "refunded_amount": amount, + }, + ) + refund_token = await send_cashu_refund( + amount, unit, mint, request_id + ) + passthrough = _passthrough_key_config_response(resp) + passthrough.headers["X-Cashu"] = refund_token + return passthrough + refund_token = await send_cashu_refund(amount, unit, mint, request_id) error_response = Response( content=json.dumps( @@ -1076,9 +1149,7 @@ async def forward_ehbp_x_cashu_request( usage_source = ( "header" if usage_header_name - and any( - k.lower() == usage_header_name.lower() for k, _ in resp.headers - ) + and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers) else ("trailer" if usage_header else "none") ) @@ -1153,6 +1224,46 @@ async def forward_ehbp_x_cashu_request( except Exception: raise + except EhbpTimeoutError as e: + logger.warning( + "EHBP X-Cashu upstream timed out", + extra={ + "error": str(e), + "path": path, + "method": request.method, + "redeemed": redeemed, + }, + ) + + if redeemed and amount > 0: + try: + refund_token = await send_cashu_refund(amount, unit, mint, request_id) + error_response = create_error_response( + "upstream_timeout", + str(e), + 504, + request=request, + code="UPSTREAM_TIMEOUT", + ) + error_response.headers["X-Cashu"] = refund_token + return error_response + except Exception as refund_error: + logger.error( + "Failed to refund EHBP X-Cashu token after timeout", + extra={ + "error": str(refund_error), + "original_error": str(e), + }, + ) + + return create_error_response( + "upstream_timeout", + str(e), + 504, + request=request, + code="UPSTREAM_TIMEOUT", + ) + except Exception as e: error_message = str(e) logger.error( @@ -1186,6 +1297,30 @@ async def forward_ehbp_x_cashu_request( }, ) + if not redeemed: + classified = classify_redemption_error(e) + if classified is not None: + error_type, status_code, message, error_code = classified + # Never re-offer a spent/consumed token. + echo_token = None if error_code in SPENT_TOKEN_CODES else x_cashu_token + return create_error_response( + error_type, + message, + status_code, + request=request, + token=echo_token, + code=error_code, + ) + # Raw exception text may contain the attacker-supplied mint URL. + return create_error_response( + "api_error", + "Internal error during token redemption", + 500, + request=request, + token=x_cashu_token, + code="internal_error", + ) + if "already spent" in error_message.lower(): return create_error_response( "token_already_spent", diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index 3d689922..c7e698e3 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -32,6 +32,7 @@ from ..core.exceptions import UpstreamError from ..core.redaction import redact_org_ids from ..payment.models import Model from .rate_limit import classify_rate_limit +from .reasoning_effort import adapt_messages_body_for_litellm logger = get_logger(__name__) @@ -40,6 +41,10 @@ logger = get_logger(__name__) # unsupported params; these newer/extension fields get passed through # verbatim and the upstream rejects them with a 400. Pop them here so the # request reaches the upstream cleanly. +# +# Note: ``dispatch_anthropic_messages`` additionally enforces +# ``ALLOWED_MESSAGES_REQUEST_FIELDS``, which already excludes all of these. +# This tuple remains for ``gemini_messages``, which pops them explicitly. ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = ( "thinking", "cache_control", @@ -51,6 +56,27 @@ ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = ( "anthropic_beta", ) +# Anthropic Messages API request fields forwarded from the client body +# into the upstream call. Only these are passed on; anything else is +# dropped so the forwarded request is deterministic and limited to the +# documented Messages surface. +ALLOWED_MESSAGES_REQUEST_FIELDS: frozenset[str] = frozenset( + { + "messages", + "max_tokens", + "system", + "temperature", + "top_p", + "top_k", + "stop_sequences", + "tools", + "tool_choice", + "metadata", + # OpenAI-shaped effort after thinking is lifted off Anthropic bodies. + "reasoning_effort", + } +) + def coerce_litellm_payload(payload: object) -> dict: """Convert a litellm event into a plain dict. @@ -108,9 +134,7 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]: return events, buffer -def events_from_chunk( - chunk: object, sse_buffer: bytes -) -> tuple[list[dict], bytes]: +def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]: """Normalize a stream chunk into one or more event dicts. ``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE @@ -201,9 +225,7 @@ async def aggregate_anthropic_events_to_message( raw_json = partial_json.pop(idx, None) if raw_json is not None and idx < len(blocks): try: - blocks[idx]["input"] = ( - json.loads(raw_json) if raw_json else {} - ) + blocks[idx]["input"] = json.loads(raw_json) if raw_json else {} except json.JSONDecodeError: blocks[idx]["input"] = raw_json elif etype == "message_delta": @@ -445,9 +467,7 @@ async def dispatch_anthropic_messages( on bad input or upstream failure. """ if not request_body: - raise UpstreamError( - "Missing request body for /v1/messages", status_code=400 - ) + raise UpstreamError("Missing request body for /v1/messages", status_code=400) try: body: dict = json.loads(request_body) @@ -466,15 +486,18 @@ async def dispatch_anthropic_messages( client_stream = bool(body.pop("stream", False)) upstream_stream = True - dropped: dict[str, Any] = {} - for field in ANTHROPIC_ONLY_FIELDS: - if field in body: - dropped[field] = body.pop(field) + adapt_messages_body_for_litellm(body, model_obj) + + # Forward only allowlisted Anthropic Messages request fields. Any + # other client-supplied key is dropped so it cannot leak into the + # upstream request. See ALLOWED_MESSAGES_REQUEST_FIELDS. + dropped = sorted(set(body) - ALLOWED_MESSAGES_REQUEST_FIELDS) if dropped: logger.debug( - "Dropped anthropic-only fields before litellm dispatch", - extra={"dropped_keys": sorted(dropped.keys())}, + "Dropped non-forwardable fields before litellm dispatch", + extra={"dropped_keys": dropped}, ) + body = {k: v for k, v in body.items() if k in ALLOWED_MESSAGES_REQUEST_FIELDS} # Convention: `model.id` is the canonical upstream model name; # `forwarded_model_id` is the public alias the internal API exposes diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 95b6edf2..092bd318 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -1,9 +1,6 @@ """Model-path discovery service. Exposes every selectable upstream route a Routstr model is reachable through. -This PR remains discovery-only: request-side routing will consume the opaque -selectors in a follow-up. - A path is a standard percent-encoded query string containing the configured upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter endpoint, its machine-readable tag:: @@ -16,11 +13,12 @@ from __future__ import annotations import asyncio import ipaddress +import json import random import time from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Callable -from urllib.parse import urlencode, urlsplit +from urllib.parse import parse_qsl, urlencode, urlsplit import httpx from sqlalchemy.dialects.sqlite import insert @@ -60,10 +58,11 @@ ModelKey = tuple[str, int] @dataclass(frozen=True) class EndpointIdentity: - """Exact OpenRouter endpoint identity returned by ``/endpoints``.""" + """Exact OpenRouter endpoint and its provider-specific model metadata.""" tag: str provider_name: str | None + model_metadata: dict[str, Any] @dataclass(frozen=True) @@ -83,6 +82,7 @@ class DiscoveredPath: model_id: str path: str provider: ConfiguredProviderIdentity + model_metadata: dict[str, Any] endpoint_tag: str | None = None endpoint_name: str | None = None @@ -132,21 +132,68 @@ def encode_model_path( return urlencode(components) +@dataclass(frozen=True) +class ModelPathSelector: + """Decoded client-supplied route selector.""" + + base_url: str + provider_id: int + model_id: str + endpoint_tag: str | None = None + + +def decode_model_path(path: str) -> ModelPathSelector | None: + """Inverse of ``encode_model_path``; ``None`` when the selector is malformed.""" + try: + pairs = parse_qsl( + path, + keep_blank_values=True, + strict_parsing=True, + max_num_fields=4, + errors="strict", + ) + except ValueError: + return None + params = dict(pairs) + if len(params) != len(pairs) or params.keys() - { + "url", + "provider-id", + "model-id", + "endpoint", + }: + return None + if any(not value.strip() for value in params.values()): + return None + base_url = params.get("url", "") + model_id = params.get("model-id", "") + raw_provider_id = params.get("provider-id", "") + if not base_url or not model_id or not raw_provider_id: + return None + try: + provider_id = int(raw_provider_id) + except ValueError: + return None + if provider_id <= 0: + return None + return ModelPathSelector( + base_url=base_url, + provider_id=provider_id, + model_id=model_id, + endpoint_tag=params.get("endpoint") or None, + ) + + def _make_http_client() -> httpx.AsyncClient: """Client factory, separated so tests can substitute a mock transport.""" return httpx.AsyncClient() def is_openrouter_base_url(base_url: str | None) -> bool: - """True when ``base_url`` points at OpenRouter. - - Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``: - that predicate also returns True for native Anthropic (correct for - cache-control, wrong for OpenRouter endpoint discovery). This one keys only - on the URL so a ``GenericUpstreamProvider`` aimed at OpenRouter is matched - while native Anthropic is not. - """ - return "openrouter.ai" in (base_url or "") + """Match OpenRouter itself, not compatible providers or lookalike hosts.""" + try: + return urlsplit(base_url or "").hostname == "openrouter.ai" + except ValueError: + return False def exposed_model_id(model: object) -> str: @@ -263,9 +310,14 @@ async def _fetch_openrouter_endpoint_subproviders( try: payload = resp.json() data = payload.get("data") if isinstance(payload, dict) else None - endpoints = data.get("endpoints") if isinstance(data, dict) else None + if not isinstance(data, dict): + raise ValueError("data must be an object") + endpoints = data.get("endpoints") if not isinstance(endpoints, list): raise ValueError("endpoints must be a list") + common_metadata = { + key: value for key, value in data.items() if key != "endpoints" + } identities: dict[str, EndpointIdentity] = {} for endpoint in endpoints: if not isinstance(endpoint, dict): @@ -281,6 +333,7 @@ async def _fetch_openrouter_endpoint_subproviders( provider_name=provider_name if isinstance(provider_name, str) and provider_name else None, + model_metadata={**common_metadata, **endpoint}, ), ) if endpoints and not identities: @@ -341,6 +394,34 @@ async def _load_model_visibility() -> tuple[ return overrides_by_key, disabled_model_keys, provider_identities +def _serialize_model_metadata(model: object, model_id: str) -> dict[str, Any]: + """Serialize provider-specific model details into the public API shape.""" + model_dict = getattr(model, "dict", None) + if callable(model_dict): + metadata = dict(model_dict()) + else: + metadata = { + key: value for key, value in vars(model).items() if not key.startswith("_") + } + + for field in ( + "architecture", + "pricing", + "sats_pricing", + "per_request_limits", + "top_provider", + "alias_ids", + ): + value = metadata.get(field) + if isinstance(value, str): + try: + metadata[field] = json.loads(value) + except (TypeError, ValueError): + pass + metadata["id"] = model_id + return metadata + + def _apply_model_visibility( upstream: BaseUpstreamProvider, overrides_by_key: dict[ModelKey, ModelRow] | None, @@ -348,11 +429,10 @@ def _apply_model_visibility( ) -> list[object]: """Return provider models after DB disabled/override state is applied. - Only the identity fields (``id``, ``forwarded_model_id``, - ``canonical_slug``) matter for path discovery, so DB override rows are used - directly rather than rebuilt into fully priced ``Model`` objects — the - pricing pipeline costs ~0.7ms of event-loop CPU per row for data this - module immediately discards. + DB override rows are used directly rather than rebuilt into priced + ``Model`` objects. Their JSON metadata fields are decoded when each path is + collected, preserving the provider-specific stored values without running + the routing price-selection pipeline. """ overrides_by_key = overrides_by_key or {} disabled_model_keys = disabled_model_keys or set() @@ -414,6 +494,7 @@ async def _collect_provider_paths( provider_identity.base_url, provider_identity.id, model_id ), provider=provider_identity, + model_metadata=_serialize_model_metadata(model, model_id), ) if not is_openrouter_base_url(upstream.base_url): @@ -453,6 +534,7 @@ async def _collect_provider_paths( endpoint.tag, ), provider=provider_identity, + model_metadata={**endpoint.model_metadata, "id": model_id}, endpoint_tag=endpoint.tag, endpoint_name=endpoint.provider_name, ) @@ -512,6 +594,7 @@ async def _persist_provider_paths( "provider_type": discovered.provider.provider_type, "endpoint_tag": discovered.endpoint_tag, "endpoint_name": discovered.endpoint_name, + "model_metadata": json.dumps(discovered.model_metadata), "upstream_provider_id": upstream_provider_id, "updated_at": now, } @@ -526,6 +609,7 @@ async def _persist_provider_paths( "provider_type": insert_stmt.excluded.provider_type, "endpoint_tag": insert_stmt.excluded.endpoint_tag, "endpoint_name": insert_stmt.excluded.endpoint_name, + "model_metadata": insert_stmt.excluded.model_metadata, "updated_at": insert_stmt.excluded.updated_at, }, ) @@ -736,10 +820,83 @@ async def refresh_model_paths_periodically( break -def _serialize_path(row: ModelPathRow) -> dict[str, Any]: +def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None: + """Run a path's USD rates through the ``/v1/models`` pricing pipeline. + + Metadata copied from the provider model cache is already priced. OpenRouter + endpoint metadata is not: it carries that endpoint's own USD rates, which + still need the cache backfill, the provider fee and the sats conversion. + """ + pricing = model.get("pricing") + if model.get("sats_pricing") or not isinstance(pricing, dict): + return + + from ..payment.models import ( + Architecture, + Model, + Pricing, + TopProvider, + _calculate_usd_max_costs, + _update_model_sats_pricing, + backfill_cache_pricing, + ) + from ..payment.price import sats_usd_price + + try: + model_id = model.get("forwarded_model_id") or model["id"] + usd = backfill_cache_pricing(model_id, Pricing.parse_obj(pricing)) + usd = Pricing.parse_obj({k: v * provider_fee for k, v in usd.dict().items()}) + priced = Model( + id=model_id, + name=model.get("name") or model_id, + created=0, + description="", + context_length=model.get("context_length") or 0, + architecture=Architecture( + modality="text", + input_modalities=[], + output_modalities=[], + tokenizer="", + instruct_type=None, + ), + pricing=usd, + top_provider=TopProvider( + context_length=model.get("context_length"), + max_completion_tokens=model.get("max_completion_tokens"), + ), + ) + ( + usd.max_prompt_cost, + usd.max_completion_cost, + usd.max_cost, + ) = _calculate_usd_max_costs(priced) + priced = _update_model_sats_pricing(priced, sats_usd_price()) + except Exception as exc: + # An endpoint with rates we cannot price is still a usable route, so it + # is served with its raw upstream pricing rather than dropped. + logger.warning( + "Could not calculate sats pricing for model path", + extra={"model_id": model.get("id"), "error": str(exc)}, + ) + return + + if priced.sats_pricing: + model["pricing"] = usd.dict() + model["sats_pricing"] = priced.sats_pricing.dict() + + +def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]: endpoint = None if row.endpoint_tag or row.endpoint_name: endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name} + try: + model = json.loads(row.model_metadata) + except (TypeError, ValueError): + model = {} + if not isinstance(model, dict): + model = {} + model.setdefault("id", row.model_id) + _price_in_sats(model, provider_fee) return { "path": row.path, "provider": { @@ -748,11 +905,17 @@ def _serialize_path(row: ModelPathRow) -> dict[str, Any]: "type": row.provider_type, }, "endpoint": endpoint, + "model": model, } +async def _provider_fees(session: "AsyncSession") -> dict[int, float]: + rows = (await session.exec(select(UpstreamProviderRow))).all() + return {row.id: row.provider_fee for row in rows if row.id is not None} + + async def get_all_model_paths() -> dict: - """All models with their exact selectable routes.""" + """All models with exact routes and provider-specific model metadata.""" async with create_session() as session: rows = ( await session.exec( @@ -763,6 +926,7 @@ async def get_all_model_paths() -> dict: ) ) ).all() + fees = await _provider_fees(session) grouped: dict[str, list[dict[str, Any]]] = {} seen_paths: dict[str, set[str]] = {} @@ -772,7 +936,9 @@ async def get_all_model_paths() -> dict: if row.path in seen_paths.setdefault(row.model_id, set()): continue seen_paths[row.model_id].add(row.path) - grouped.setdefault(row.model_id, []).append(_serialize_path(row)) + grouped.setdefault(row.model_id, []).append( + _serialize_path(row, fees.get(row.upstream_provider_id, 1.01)) + ) data = [ { "id": grouped_model_id, @@ -806,6 +972,7 @@ async def get_paths_for_model(model_id: str) -> dict: unprefixed_id = public_model_id(model_id) if unprefixed_id != model_id: rows = await load_rows(session, unprefixed_id) + fees = await _provider_fees(session) seen: set[str] = set() paths: list[dict] = [] @@ -815,5 +982,5 @@ async def get_paths_for_model(model_id: str) -> dict: if row.path in seen: continue seen.add(row.path) - paths.append(_serialize_path(row)) + paths.append(_serialize_path(row, fees.get(row.upstream_provider_id, 1.01))) return {"data": paths, "updated_at": updated_at or None} diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index 9fed0154..46752169 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -66,9 +66,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): """Strip 'ollama/' prefix for Ollama API compatibility.""" return model_id.removeprefix("ollama/") - def get_request_base_url( - self, path: str, model_obj: Model | None = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str: """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint.""" return f"{self.base_url.rstrip('/')}/v1" @@ -185,7 +183,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): except Exception: self._models_cache = models_with_fees - self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache} + self._models_by_id = { + m.forwarded_model_id or m.id: m for m in self._models_cache + } logger.info( f"Refreshed models cache for {self.base_url}", extra={"model_count": len(models)}, @@ -224,26 +224,14 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): Returns: Model with provider fee applied to pricing and max costs calculated """ - from ..payment.models import Model, Pricing, _calculate_usd_max_costs + from ..payment.models import Pricing, _calculate_usd_max_costs adjusted_pricing = Pricing.parse_obj( {k: v * self.provider_fee for k, v in model.pricing.dict().items()} ) - temp_model = Model( - id=model.id, - name=model.name, - created=model.created, - description=model.description, - context_length=model.context_length, - architecture=model.architecture, - pricing=adjusted_pricing, - sats_pricing=None, - per_request_limits=model.per_request_limits, - top_provider=model.top_provider, - enabled=model.enabled, - upstream_provider_id=model.upstream_provider_id, - canonical_slug=model.canonical_slug, + temp_model = model.copy( + update={"pricing": adjusted_pricing, "sats_pricing": None} ) ( @@ -252,18 +240,4 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): adjusted_pricing.max_cost, ) = _calculate_usd_max_costs(temp_model) - return Model( - id=model.id, - name=model.name, - created=model.created, - description=model.description, - context_length=model.context_length, - architecture=model.architecture, - pricing=adjusted_pricing, - sats_pricing=model.sats_pricing, - per_request_limits=model.per_request_limits, - top_provider=model.top_provider, - enabled=model.enabled, - upstream_provider_id=model.upstream_provider_id, - canonical_slug=model.canonical_slug, - ) + return model.copy(update={"pricing": adjusted_pricing}) diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 50ab4532..65ccbb57 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -1,5 +1,9 @@ from __future__ import annotations +import asyncio +import random +import time +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Optional import httpx @@ -15,6 +19,87 @@ if TYPE_CHECKING: logger = get_logger(__name__) +_PPQ_SAFE_READ_ATTEMPTS = 3 +_PPQ_CIRCUIT_COOLDOWN_SECONDS = 30.0 + + +class PPQCircuitOpenError(RuntimeError): + pass + + +@dataclass +class _PPQCircuitState: + consecutive_failures: int = 0 + cooldown_until: float = 0.0 + lock: asyncio.Lock = field(default_factory=asyncio.Lock) + loop: asyncio.AbstractEventLoop | None = None + + +_ppq_circuits: dict[str, _PPQCircuitState] = {} + + +def _ppq_origin(url: str) -> str: + parsed = httpx.URL(url) + port = parsed.port or {"https": 443, "http": 80}.get(parsed.scheme, 0) + return f"{parsed.scheme}://{parsed.host}:{port}" + + +async def _safe_read_request( + client: httpx.AsyncClient, + method: str, + url: str, + *, + headers: dict[str, str], + json: dict[str, object] | None = None, +) -> httpx.Response: + state = _ppq_circuits.setdefault(_ppq_origin(url), _PPQCircuitState()) + loop = asyncio.get_running_loop() + if state.loop is not loop: + # Locks cannot be reused across event loops. + state.lock = asyncio.Lock() + state.loop = loop + async with state.lock: + remaining = state.cooldown_until - time.monotonic() + if remaining > 0: + raise PPQCircuitOpenError( + f"PPQ.AI safe-read circuit is open; retry after {remaining:.2f}s" + ) + + for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1): + try: + response = await client.request(method, url, headers=headers, json=json) + response.raise_for_status() + state.consecutive_failures = 0 + state.cooldown_until = 0.0 + return response + except (httpx.TransportError, httpx.HTTPStatusError) as error: + retryable_status = isinstance(error, httpx.HTTPStatusError) and ( + error.response.status_code in {502, 503, 504} + ) + if not isinstance(error, httpx.TransportError) and not retryable_status: + raise + state.consecutive_failures += 1 + if attempt >= _PPQ_SAFE_READ_ATTEMPTS: + state.cooldown_until = ( + time.monotonic() + _PPQ_CIRCUIT_COOLDOWN_SECONDS + ) + raise + base_delay = 0.25 * (2 ** (attempt - 1)) + delay = base_delay + random.uniform(0.0, base_delay) + logger.warning( + "PPQ.AI safe read failed; retrying", + extra={ + "url": url, + "attempt": attempt, + "max_attempts": _PPQ_SAFE_READ_ATTEMPTS, + "backoff_seconds": round(delay, 3), + "error": repr(error), + "error_type": type(error).__name__, + }, + ) + await asyncio.sleep(delay) + raise RuntimeError("unreachable") + class PPQAIModelPricing(BaseModel): ui: Optional[dict[str, float]] = None @@ -123,122 +208,109 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): url = f"{self.base_url}/models" headers = {"Authorization": f"Bearer {self.api_key}"} - try: - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() - data = response.json() + async with httpx.AsyncClient(timeout=30.0) as client: + response = await _safe_read_request(client, "GET", url, headers=headers) + data = response.json() - models_data = data.get("data", []) + models_data = data.get("data", []) - or_models = [ - Model(**model) # type: ignore - for model in await async_fetch_openrouter_models() - ] + or_models = [ + Model(**model) # type: ignore + for model in await async_fetch_openrouter_models() + ] - models = [] - for model_data in models_data: - try: - ppqai_model = PPQAIModel.parse_obj(model_data) - if ppqai_model.id in self.IGNORED_MODEL_IDS: - continue + models = [] + for model_data in models_data: + try: + ppqai_model = PPQAIModel.parse_obj(model_data) + if ppqai_model.id in self.IGNORED_MODEL_IDS: + continue - or_model = next( - ( - model - for model in or_models - if (model.id == ppqai_model.id) - or (model.id.split("/")[-1] == ppqai_model.id) - or (model.id == ppqai_model.id.split("/")[-1]) - ), - None, - ) + or_model = next( + ( + model + for model in or_models + if (model.id == ppqai_model.id) + or (model.id.split("/")[-1] == ppqai_model.id) + or (model.id == ppqai_model.id.split("/")[-1]) + ), + None, + ) - if or_model: - input_price = None - if ppqai_model.pricing.api: - input_price = ppqai_model.pricing.api.get( - "input_per_1M" - ) - elif ppqai_model.pricing.input_per_1M_tokens: - input_price = ppqai_model.pricing.input_per_1M_tokens + if or_model: + input_price = None + if ppqai_model.pricing.api: + input_price = ppqai_model.pricing.api.get("input_per_1M") + elif ppqai_model.pricing.input_per_1M_tokens: + input_price = ppqai_model.pricing.input_per_1M_tokens - if input_price is not None: - or_model.pricing.prompt = input_price / 1_000_000 + if input_price is not None: + or_model.pricing.prompt = input_price / 1_000_000 - output_price = None - if ppqai_model.pricing.api: - output_price = ppqai_model.pricing.api.get( - "output_per_1M" - ) - elif ppqai_model.pricing.output_per_1M_tokens: - output_price = ppqai_model.pricing.output_per_1M_tokens + output_price = None + if ppqai_model.pricing.api: + output_price = ppqai_model.pricing.api.get("output_per_1M") + elif ppqai_model.pricing.output_per_1M_tokens: + output_price = ppqai_model.pricing.output_per_1M_tokens - if output_price is not None: - or_model.pricing.completion = output_price / 1_000_000 + if output_price is not None: + or_model.pricing.completion = output_price / 1_000_000 - if cl := ppqai_model.context_length: - or_model.context_length = cl - models.append(or_model) - else: - input_price = 0.0 - if ppqai_model.pricing.api: - input_price = ppqai_model.pricing.api.get( - "input_per_1M", 0.0 - ) - elif ppqai_model.pricing.input_per_1M_tokens: - input_price = ppqai_model.pricing.input_per_1M_tokens - - output_price = 0.0 - if ppqai_model.pricing.api: - output_price = ppqai_model.pricing.api.get( - "output_per_1M", 0.0 - ) - elif ppqai_model.pricing.output_per_1M_tokens: - output_price = ppqai_model.pricing.output_per_1M_tokens - - models.append( - Model( - id=ppqai_model.id, - name=ppqai_model.name, - created=ppqai_model.created_at // 1000, - description=f"{ppqai_model.provider or 'PPQ.AI'} model", - context_length=ppqai_model.context_length, - architecture=Architecture( - modality="text->text", - input_modalities=["text"], - output_modalities=["text"], - tokenizer="Unknown", - instruct_type=None, - ), - pricing=Pricing( - prompt=input_price / 1_000_000, - completion=output_price / 1_000_000, - request=0.0, - image=0.0, - web_search=0.0, - internal_reasoning=0.0, - ), - ) + if cl := ppqai_model.context_length: + or_model.context_length = cl + models.append(or_model) + else: + input_price = 0.0 + if ppqai_model.pricing.api: + input_price = ppqai_model.pricing.api.get( + "input_per_1M", 0.0 + ) + elif ppqai_model.pricing.input_per_1M_tokens: + input_price = ppqai_model.pricing.input_per_1M_tokens + + output_price = 0.0 + if ppqai_model.pricing.api: + output_price = ppqai_model.pricing.api.get( + "output_per_1M", 0.0 + ) + elif ppqai_model.pricing.output_per_1M_tokens: + output_price = ppqai_model.pricing.output_per_1M_tokens + + models.append( + Model( + id=ppqai_model.id, + name=ppqai_model.name, + created=ppqai_model.created_at // 1000, + description=f"{ppqai_model.provider or 'PPQ.AI'} model", + context_length=ppqai_model.context_length, + architecture=Architecture( + modality="text->text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="Unknown", + instruct_type=None, + ), + pricing=Pricing( + prompt=input_price / 1_000_000, + completion=output_price / 1_000_000, + request=0.0, + image=0.0, + web_search=0.0, + internal_reasoning=0.0, + ), ) - except Exception as e: - logger.warning( - "Failed to parse PPQ.AI model", - extra={ - "model_id": model_data.get("id", "unknown"), - "error": str(e), - "error_type": type(e).__name__, - }, ) + except Exception as e: + logger.warning( + "Failed to parse PPQ.AI model", + extra={ + "model_id": model_data.get("id", "unknown"), + "error": str(e), + "error_type": type(e).__name__, + }, + ) - return models - - except Exception as e: - logger.error( - "Error fetching models from PPQ.AI", - extra={"error": str(e), "error_type": type(e).__name__}, - ) - return [] + return models async def on_upstream_error_redirect( self, status_code: int, error_message: str @@ -360,8 +432,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): ) async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() + response = await _safe_read_request(client, "GET", url, headers=headers) status_data = response.json() is_paid = status_data.get("status") == "Settled" @@ -460,8 +531,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): logger.debug("Checking PPQ.AI account balance", extra={"url": url}) async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.post(url, headers=headers, json={}) - response.raise_for_status() + response = await _safe_read_request( + client, "POST", url, headers=headers, json={} + ) balance_data = response.json() logger.debug( diff --git a/routstr/upstream/pricing_resolver.py b/routstr/upstream/pricing_resolver.py index 3d009fdc..8b58dddb 100644 --- a/routstr/upstream/pricing_resolver.py +++ b/routstr/upstream/pricing_resolver.py @@ -17,6 +17,8 @@ from __future__ import annotations from dataclasses import dataclass, field +from ..payment.rates import coerce_rate + @dataclass class ResolvedPricing: @@ -65,11 +67,8 @@ def estimate_context_length(model_id: str) -> int: def _as_float(value: object) -> float | None: - """OpenRouter reports prices as strings; coerce, ``None`` if unparseable.""" - try: - return float(value) # type: ignore[arg-type] - except (TypeError, ValueError): - return None + """OpenRouter reports prices as strings; coerce, ``None`` if not a real rate.""" + return coerce_rate(value) def _as_int(value: object) -> int | None: @@ -78,23 +77,24 @@ def _as_int(value: object) -> int | None: def _from_litellm(model_id: str) -> ResolvedPricing | None: - # Lazy import so the resolver stays import-light and shares the exact - # lookup semantics used by cache-rate backfill. + # Lazy import so the resolver shares the exact lookup semantics used by + # cache-rate backfill without importing the models module at load time. from ..payment.models import litellm_cost_entry info = litellm_cost_entry(model_id) if info is None: return None - prompt = info.get("input_cost_per_token") - completion = info.get("output_cost_per_token") - if not isinstance(prompt, (int, float)) or not isinstance(completion, (int, float)): + prompt = coerce_rate(info.get("input_cost_per_token")) + completion = coerce_rate(info.get("output_cost_per_token")) + if prompt is None or completion is None: return None # A both-zero entry is litellm listing a model without a real price (free # moderation/rerank tiers do this) — treating 0/0 as resolved would serve - # the model for free. Reject it (and any negative) so the caller falls - # through, mirroring async_fetch_openrouter_models' _has_valid_pricing. - if prompt < 0 or completion < 0 or (prompt == 0 and completion == 0): + # the model for free. Reject it so the caller falls through, mirroring + # async_fetch_openrouter_models' _has_valid_pricing. Coercion runs first: + # `NaN` would defeat this guard on its own, every comparison being False. + if prompt == 0 and completion == 0: return None input_modalities = ["text"] @@ -102,8 +102,8 @@ def _from_litellm(model_id: str) -> ResolvedPricing | None: input_modalities.append("image") return ResolvedPricing( - prompt=float(prompt), - completion=float(completion), + prompt=prompt, + completion=completion, # max_input_tokens is the context window; max_tokens is litellm's # completion cap (it tracks max_output_tokens for ~94% of models), so # it is never a context source. A missing window falls to the id-based diff --git a/routstr/upstream/reasoning_effort.py b/routstr/upstream/reasoning_effort.py new file mode 100644 index 00000000..994135ad --- /dev/null +++ b/routstr/upstream/reasoning_effort.py @@ -0,0 +1,197 @@ +"""Map client reasoning/thinking effort onto a model's allowlist. + +OpenRouter (and some other catalogs) publish per-model reasoning metadata: +which effort levels are legal, the default, and whether reasoning is +mandatory. Clients still send the generic OpenAI / Anthropic shapes +(``reasoning_effort``, ``reasoning.effort``, ``thinking``). This module +normalizes those into a supported effort and writes the fields the +upstream actually accepts, instead of dropping the parameter or +forwarding a value the model rejects. +""" + +from __future__ import annotations + +from typing import Any + +from ..payment.models import Model, Reasoning + +# Highest first. Unknown values are treated as unranked. +EFFORT_RANK: tuple[str, ...] = ( + "max", + "xhigh", + "high", + "medium", + "low", + "minimal", + "none", +) +_RANK_INDEX: dict[str, int] = {name: i for i, name in enumerate(EFFORT_RANK)} + +_REASONING_KEYS = ("reasoning", "reasoning_effort", "thinking") + + +def _normalize_effort(value: object) -> str | None: + if not isinstance(value, str): + return None + cleaned = value.strip().lower() + return cleaned or None + + +def closest_supported_effort( + requested: str | None, + supported: list[str], + *, + default_effort: str | None = None, + mandatory: bool = False, +) -> str | None: + """Pick a legal effort for ``requested``. + + Exact match wins. Otherwise the nearest rank in ``EFFORT_RANK`` is + used (preferring the higher neighbour on a tie). ``none`` is rejected + when ``mandatory`` is set. Missing / unmapped requests fall back to + ``default_effort``, then the highest remaining supported level. + """ + allowed_efforts: list[str] = [ + normalized + for item in supported + if (normalized := _normalize_effort(item)) is not None + ] + if mandatory: + allowed_efforts = [item for item in allowed_efforts if item != "none"] + if not allowed_efforts: + if mandatory: + return _normalize_effort(default_effort) + return _normalize_effort(requested) or _normalize_effort(default_effort) + + default = _normalize_effort(default_effort) + if default not in allowed_efforts: + default = allowed_efforts[0] + + requested_norm = _normalize_effort(requested) + if requested_norm is None or (requested_norm == "none" and mandatory): + return default + + if requested_norm in allowed_efforts: + return requested_norm + + if requested_norm not in _RANK_INDEX: + return default + + target = _RANK_INDEX[requested_norm] + return min( + allowed_efforts, + key=lambda effort: ( + abs(_RANK_INDEX.get(effort, 10_000) - target), + _RANK_INDEX.get(effort, 10_000), + ), + ) + + +def resolve_effort(requested: str | None, reasoning: Reasoning | None) -> str | None: + """Map ``requested`` through ``reasoning`` metadata when present.""" + if reasoning is None: + return _normalize_effort(requested) + supported = reasoning.supported_efforts or [] + return closest_supported_effort( + requested, + supported, + default_effort=reasoning.default_effort, + mandatory=bool(reasoning.mandatory), + ) + + +def _effort_from_thinking(thinking: object) -> str | None: + if not isinstance(thinking, dict): + return None + effort = _normalize_effort(thinking.get("effort")) + if effort: + return effort + thinking_type = _normalize_effort(thinking.get("type")) + if thinking_type in {"disabled", "none"}: + return "none" + return None + + +def extract_requested_effort(data: dict[str, Any]) -> str | None: + """Best-effort effort string from the OpenAI / Anthropic request shapes.""" + if isinstance(data.get("reasoning"), dict): + nested = _normalize_effort(data["reasoning"].get("effort")) + if nested: + return nested + top_level = _normalize_effort(data.get("reasoning_effort")) + if top_level: + return top_level + return _effort_from_thinking(data.get("thinking")) + + +def _request_mentions_reasoning(data: dict[str, Any]) -> bool: + return any(key in data for key in _REASONING_KEYS) + + +def apply_reasoning_effort( + data: dict[str, Any], + model: Model, + *, + drop_thinking: bool = False, +) -> bool: + """Rewrite ``data`` in place so effort matches the model allowlist. + + Returns True when ``data`` changed. Leaves the body alone when the + caller did not send a reasoning field and the model does not require + one. Existing ``reasoning`` object keys (``max_tokens``, ``exclude``, + ``enabled``) are preserved; only ``effort`` is mapped. + + ``drop_thinking`` is for OpenAI-compatible backends that reject the + Anthropic ``thinking`` object: the effort is lifted onto + ``reasoning_effort`` / ``reasoning.effort`` and ``thinking`` is removed. + """ + if not isinstance(data, dict): + return False + + reasoning_meta = getattr(model, "reasoning", None) + mentioned = _request_mentions_reasoning(data) + if not mentioned and not (reasoning_meta and reasoning_meta.mandatory): + return False + + resolved = resolve_effort(extract_requested_effort(data), reasoning_meta) + changed = False + + if drop_thinking and "thinking" in data: + data.pop("thinking", None) + changed = True + + if resolved is None: + return changed + + if "reasoning_effort" in data: + if data.get("reasoning_effort") != resolved: + data["reasoning_effort"] = resolved + changed = True + elif drop_thinking or (reasoning_meta and reasoning_meta.mandatory): + # Invent the OpenAI-shaped field when we stripped Anthropic + # ``thinking``, or when the model will reject a request with no + # effort at all. + if not isinstance(data.get("reasoning"), dict): + data["reasoning_effort"] = resolved + changed = True + + existing = data.get("reasoning") + if isinstance(existing, dict): + if existing.get("effort") != resolved: + data["reasoning"] = {**existing, "effort": resolved} + changed = True + elif reasoning_meta and reasoning_meta.mandatory and "reasoning_effort" not in data: + data["reasoning"] = {"effort": resolved} + changed = True + + return changed + + +def adapt_messages_body_for_litellm(data: dict[str, Any], model: Model) -> None: + """Convert Anthropic ``thinking`` into OpenAI-shaped effort for litellm. + + Litellm's Anthropic-messages adapter talking to an OpenAI-compatible + upstream will 400 on ``thinking``. Lift the effort onto + ``reasoning_effort`` and drop the Anthropic-only object. + """ + apply_reasoning_effort(data, model, drop_thinking=True) diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py index c2ea5b1d..d959a015 100644 --- a/routstr/upstream/request_correction.py +++ b/routstr/upstream/request_correction.py @@ -84,13 +84,33 @@ def extract_error_message(response: Response) -> str: return "" -def strip_unsupported_param( - body: dict, error_message: str -) -> tuple[dict, str] | None: +# Spend-shaping fields bound how much work — and therefore cost — the upstream +# may perform. The reservation was priced with these caps in place; dropping one +# and retrying would let the request run uncapped (or fan out) and bill above the +# caller's authorization. When the upstream names one of these, decline the strip +# and let the error propagate. Matched case-insensitively. +_SPEND_SHAPING_PARAMS = frozenset( + { + "max_tokens", + "max_completion_tokens", + "max_output_tokens", + "max_tokens_to_sample", + "n", + "best_of", + } +) + + +def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: """Drop a top-level param the upstream named as unsupported/deprecated. Returns ``(new_body, param)`` (a new dict, original untouched) when the error names a top-level param present in the body, otherwise ``None``. + + Spend-shaping fields (output caps, fan-out counts) are never stripped: + removing one after the reservation was priced would uncap the retry and + overcharge. When the upstream names such a field, decline so the original + error propagates rather than silently resizing the request's cost. """ match = _UNSUPPORTED_PARAM_RE.search(error_message) if not match: @@ -98,6 +118,13 @@ def strip_unsupported_param( param = match.group("param") if param not in body: return None + if param.lower() in _SPEND_SHAPING_PARAMS: + logger.warning( + "Upstream rejected spend-shaping param '%s'; refusing to strip it " + "(retrying uncapped would overcharge) — surfacing the error", + param, + ) + return None new_body = {k: v for k, v in body.items() if k != param} return new_body, param diff --git a/routstr/upstream/tinfoil.py b/routstr/upstream/tinfoil.py index 6eeb1147..0928cbc9 100644 --- a/routstr/upstream/tinfoil.py +++ b/routstr/upstream/tinfoil.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Optional import httpx from fastapi import Request @@ -28,6 +28,7 @@ logger = get_logger(__name__) class TinfoilModelPricing(BaseModel): inputTokenPricePer1M: float = 0.0 outputTokenPricePer1M: float = 0.0 + cachedInputTokenPricePer1M: Optional[float] = None requestPrice: float = 0.0 @@ -186,6 +187,14 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider): output_price = tf.pricing.outputTokenPricePer1M request_price = tf.pricing.requestPrice + # Tinfoil bills cache reads at the cached rate when the + # model exposes one, otherwise at the full input rate. + # Cache writes are never priced separately — a miss is + # just regular input prefill. + cached_price = tf.pricing.cachedInputTokenPricePer1M + if cached_price is None or cached_price <= 0.0: + cached_price = input_price + modality = "text->text" input_modalities = ["text"] output_modalities = ["text"] @@ -214,6 +223,8 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider): image=0.0, web_search=0.0, internal_reasoning=0.0, + input_cache_read=cached_price / 1_000_000, + input_cache_write=input_price / 1_000_000, ), ) ) diff --git a/routstr/upstream/tinfoil_trailer.py b/routstr/upstream/tinfoil_trailer.py index 0357864f..ec3065d6 100644 --- a/routstr/upstream/tinfoil_trailer.py +++ b/routstr/upstream/tinfoil_trailer.py @@ -20,11 +20,12 @@ from urllib.parse import urlsplit import h11 from ..core import get_logger +from ..core.exceptions import EhbpTimeoutError logger = get_logger(__name__) _READ_BUFSIZE = 65536 -_DEFAULT_TIMEOUT_SECONDS = 30.0 +_DEFAULT_TIMEOUT_SECONDS = 60.0 _DEFAULT_CLOSE_TIMEOUT_SECONDS = 1.0 _DEFAULT_MAX_RESPONSE_BYTES = 25 * 1024 * 1024 _HOP_BY_HOP_HEADERS = { @@ -100,10 +101,15 @@ async def forward_with_trailer( headers = _strip_hop_by_hop_headers(headers) ssl_ctx = ssl.create_default_context() - reader, writer = await asyncio.wait_for( - asyncio.open_connection(host, port, ssl=ssl_ctx), - timeout=timeout_seconds, - ) + try: + reader, writer = await asyncio.wait_for( + asyncio.open_connection(host, port, ssl=ssl_ctx), + timeout=timeout_seconds, + ) + except asyncio.TimeoutError as exc: + raise EhbpTimeoutError( + f"EHBP upstream {host} timed out after {timeout_seconds:g}s connecting" + ) from exc try: # Build HTTP/1.1 request @@ -126,7 +132,13 @@ async def forward_with_trailer( request_data += body writer.write(request_data) - await asyncio.wait_for(writer.drain(), timeout=timeout_seconds) + try: + await asyncio.wait_for(writer.drain(), timeout=timeout_seconds) + except asyncio.TimeoutError as exc: + raise EhbpTimeoutError( + f"EHBP upstream {host} timed out after " + f"{timeout_seconds:g}s sending request" + ) from exc # Parse response with h11 conn = h11.Connection(h11.CLIENT) @@ -140,10 +152,16 @@ async def forward_with_trailer( event = conn.next_event() if event is h11.NEED_DATA: - data = await asyncio.wait_for( - reader.read(_READ_BUFSIZE), - timeout=timeout_seconds, - ) + try: + data = await asyncio.wait_for( + reader.read(_READ_BUFSIZE), + timeout=timeout_seconds, + ) + except asyncio.TimeoutError as exc: + raise EhbpTimeoutError( + f"EHBP upstream {host} timed out after " + f"{timeout_seconds:g}s waiting for response data" + ) from exc conn.receive_data(data if data else b"") continue diff --git a/routstr/wallet.py b/routstr/wallet.py index f83ff8a6..beecde99 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -2,28 +2,29 @@ import asyncio import fcntl import json import os -import re import time import typing from contextlib import asynccontextmanager from contextvars import ContextVar from dataclasses import dataclass from pathlib import Path -from typing import AsyncGenerator, TypedDict +from typing import AsyncGenerator, Awaitable, Callable, TypedDict +from urllib.parse import urlsplit, urlunsplit import httpx -from cashu.core.base import MeltQuote, MeltQuoteState, MintQuote, Proof, Token +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.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 .core import db, get_logger from .core.db import store_cashu_transaction_with_retry as store_cashu_transaction from .core.settings import settings from .mint import ( - MINT_TRANSPORT_COOLDOWN_SECONDS, MINT_TRANSPORT_EXCEPTIONS, MintError, MintRateGuard, @@ -36,6 +37,10 @@ from .mint import ( ) from .payment.lnurl import 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. +install_cashu_httpx_shim() + # Backwards-compatible aliases for callers/tests that imported the former # wallet-local policy. Production modules use the public routstr.mint API. _MintRateGuard = MintRateGuard @@ -121,6 +126,12 @@ def _mints_to_inspect() -> list[str]: return mint_urls +_WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS = 30 +_WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS = 300 +_mint_metadata_last_load: dict[str, float] = {} +_mint_metadata_load_locks: dict[str, asyncio.Lock] = {} + + class Wallet(_CashuWallet): """Cashu adapter that preserves HTTP 429 for Routstr's mint policy.""" @@ -141,11 +152,34 @@ class Wallet(_CashuWallet): _CashuWallet.raise_on_error_request(resp) async def load_mint( - self, keyset_id: str = "", force_old_keysets: bool = False + self, + keyset_id: str = "", + force_old_keysets: bool = False, + *, + force_refresh: bool = False, ) -> None: - await self.load_mint_keysets(force_old_keysets) - await self.activate_keyset(keyset_id) - await self.load_mint_info(reload=True) + mint_url = str(self.url) + lock = _mint_metadata_load_locks.setdefault(mint_url, asyncio.Lock()) + async with lock: + now = time.monotonic() + last = _mint_metadata_last_load.get(mint_url) + if ( + not force_refresh + and last is not None + and now - last < _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS + ): + try: + await self.load_keysets_from_db() + await self.activate_keyset(keyset_id) + await self.load_mint_info(reload=False) + return + except Exception: + pass + + await self.load_mint_keysets(force_old_keysets) + await self.activate_keyset(keyset_id) + await self.load_mint_info(reload=True) + _mint_metadata_last_load[mint_url] = time.monotonic() class MintConnectionError(Exception): @@ -159,6 +193,10 @@ class SourceMintConnectionError(MintConnectionError): """The mint that issued the incoming proofs cannot be reached.""" +class UntrustedSourceMintError(ValueError): + """The token names a mint outside primary_mint/cashu_mints.""" + + class TokenConsumedError(Exception): """A failure that happened AFTER the token's proofs were spent (melt succeeded, or redemption already returned) — e.g. minting on the primary @@ -170,6 +208,31 @@ class TokenConsumedError(Exception): """ +_MINT_TIMEOUT_EXCEPTIONS: tuple[type[BaseException], ...] = ( + httpx.TimeoutException, + asyncio.TimeoutError, +) + + +def _exception_chain(error: BaseException) -> typing.Iterator[BaseException]: + seen: set[int] = set() + current: BaseException | None = error + while current is not None and id(current) not in seen: + seen.add(id(current)) + yield current + current = current.__cause__ or current.__context__ + + +def is_mint_timeout(error: BaseException) -> bool: + """True if the mint accepted the connection but did not answer in time.""" + for current in _exception_chain(error): + if isinstance(current, TokenConsumedError): + return False + if isinstance(current, _MINT_TIMEOUT_EXCEPTIONS): + return True + return False + + def is_source_mint_connection_error(error: BaseException) -> bool: seen: set[int] = set() current: BaseException | None = error @@ -233,25 +296,39 @@ def classify_redemption_error( "Token was redeemed but could not be credited; do not retry", "cashu_token_consumed", ) - if is_source_mint_connection_error(error): + if isinstance(error, UntrustedSourceMintError): return ( - "mint_unreachable", - 503, - "The mint that issued this Cashu token is unreachable; the token cannot be redeemed at another mint", - "cashu_source_mint_unreachable", + "untrusted_mint", + 400, + "Cashu token was issued by a mint this node does not accept", + "cashu_untrusted_source_mint", ) if is_mint_rate_limited(error): return ( "mint_rate_limited", 503, - "Cashu mint rate-limited; retry after cooldown", + "Cashu mint is rate-limiting requests; retry later", "cashu_mint_rate_limited", ) + if is_mint_timeout(error): + return ( + "mint_timeout", + 503, + "Cashu mint did not respond in time; retry later", + "cashu_mint_timeout", + ) + if is_source_mint_connection_error(error): + return ( + "mint_unreachable", + 503, + "The mint that issued this Cashu token is unreachable; retry later", + "cashu_source_mint_unreachable", + ) if is_mint_connection_error(error): return ( "mint_unreachable", 503, - "Cashu mint is unreachable", + "Cashu mint is unreachable; retry later", "cashu_mint_unreachable", ) lowered = str(error).lower() @@ -274,13 +351,6 @@ def classify_redemption_error( "Token value is too small to cover swap fees", "cashu_token_swap_fees_exceed_amount", ) - if "failed to melt" in lowered: - return ( - "mint_error", - 422, - "Failed to swap token from foreign mint", - "cashu_foreign_mint_swap_failed", - ) if ("invalid" in lowered or "decode" in lowered) and "token" in lowered: # Anchored to "token" so internal faults whose text merely contains # "invalid"/"decode" fall through to the 500 branch, not a token error. @@ -424,64 +494,81 @@ async def _redeem_same_mint( async def recieve_token( token: str, - destination_mint: str | None = None, destination_unit: str | None = None, ) -> tuple[int, str, str]: # amount, unit, mint_url - """Redeem a token while serializing all wallet proof mutation.""" + """Redeem a token on its own (trusted) mint while serializing proof mutation.""" async with wallet_operation_guard(): - return await _recieve_token_locked(token, destination_mint, destination_unit) + return await _recieve_token_locked(token, destination_unit) + + +def _normalized_mint_url(mint_url: str) -> str: + """Fold cosmetic URL differences for trust comparison; keep path case, + userinfo, non-default port, query and fragment so hosts never alias.""" + stripped = mint_url.strip() + if any(char in stripped for char in "\t\r\n"): + return stripped + try: + parts = urlsplit(stripped) + port = parts.port + except ValueError: + return stripped + scheme = parts.scheme.lower() + if (scheme == "https" and port == 443) or (scheme == "http" and port == 80): + port = None + netloc = (parts.hostname or "").lower() + if parts.username is not None or parts.password is not None: + netloc = f"{parts.username or ''}:{parts.password or ''}@{netloc}" + if port is not None: + netloc = f"{netloc}:{port}" + return urlunsplit( + (scheme, netloc, parts.path.rstrip("/"), parts.query, parts.fragment) + ) + + +def resolve_trusted_source_mint(mint_url: str) -> str | None: + """Return the operator's spelling of ``mint_url`` if it is a trusted mint.""" + normalized = _normalized_mint_url(mint_url) + if not normalized: + return None + for candidate in _trusted_destination_candidates(): + if candidate and _normalized_mint_url(candidate) == normalized: + return candidate + return None + + +def is_trusted_source_mint(mint_url: str) -> bool: + """True if ``mint_url`` is the primary mint or one of the configured mints.""" + return resolve_trusted_source_mint(mint_url) is not None async def _recieve_token_locked( token: str, - destination_mint: str | None = None, destination_unit: str | None = None, ) -> tuple[int, str, str]: + """Redeem on the issuing mint; only trusted mints are accepted, never swapped.""" token_obj = deserialize_token_from_string(token) - if len(token_obj.keysets) > 1: - raise ValueError("Multiple keysets per token currently not supported") - - destinations = ( - [destination_mint] - if destination_mint is not None - else list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) - ) - output_unit = ( - token_obj.unit if token_obj.mint in destinations else settings.primary_mint_unit - ) - if destination_unit is not None and output_unit != destination_unit: + mint_url = resolve_trusted_source_mint(token_obj.mint) + if mint_url is None: + raise UntrustedSourceMintError(f"Untrusted source mint: {token_obj.mint}") + if destination_unit is not None and token_obj.unit != destination_unit: raise ValueError( "Cashu token unit does not match the API key liability unit: " - f"expected {destination_unit}, got {output_unit}" - ) - - wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) - if token_obj.mint not in destinations: - logger.info( - "Cashu cross-mint swap required", - extra={ - "event": "cashu_swap_started", - "source_mint": token_obj.mint, - "source_unit": token_obj.unit, - "source_amount": token_obj.amount, - "destination_candidates": destinations, - }, - ) - return await swap_to_trusted_mint( - token_obj, wallet, destination_mints=destinations + f"expected {destination_unit}, got {token_obj.unit}" ) + wallet = await get_wallet(mint_url, token_obj.unit, load=False) logger.info( "Trying same-mint Cashu redemption", extra={ "event": "cashu_same_mint_redemption", - "source_mint": token_obj.mint, + "source_mint": mint_url, "source_unit": token_obj.unit, "source_amount": token_obj.amount, "cross_mint_fallback_on_connection_failure": False, }, ) - return await _redeem_same_mint(wallet, token_obj) + amount, unit, _ = await _redeem_same_mint(wallet, token_obj) + return amount, unit, mint_url async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: @@ -491,7 +578,11 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int async def _send_locked( - amount: int, unit: str, mint_url: str | None = None + amount: int, + unit: str, + mint_url: str | None = None, + *, + owner_only: bool = False, ) -> tuple[int, str]: effective_mint_url = await find_trusted_mint_with_funds( amount, unit, mint_url, force_reload=True @@ -501,6 +592,12 @@ async def _send_locked( wallet, effective_mint_url, unit, not_reserved=True ) proofs_for_mint = sum(proof.amount for proof in proofs) + if owner_only: + owner_balance = await _owner_balance_for_mint_and_unit( + effective_mint_url, unit, proofs_for_mint + ) + if owner_balance < amount: + raise ValueError("Owner Cashu balance is insufficient for auto top-up") all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved) @@ -547,6 +644,13 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str return token +async def send_token_from_owner_locked( + amount: int, unit: str, mint_url: str | None = None +) -> str: + _, token = await _send_locked(amount, unit, mint_url, owner_only=True) + return token + + class Bolt11PaymentNotAttempted(Exception): """The invoice was definitively not paid, so the attempt can be retried. @@ -891,83 +995,6 @@ async def find_trusted_mint_with_funds( ) -# A foreign mint's fee_reserve is a non-binding estimate (NUT-05): the mint may -# demand more when re-quoting or at melt execution. Instead of padding the -# estimate with a safety buffer (which strands the margin at the foreign mint -# on every swap), the swap retries with the amount recomputed from the fees the -# mint actually demands, up to this many attempts. -_MAX_SWAP_ATTEMPTS = 3 - -_MINT_ERROR_CODE_RE = re.compile(r"\(Code: (\d+)\)") -_MELT_SHORTFALL_RE = re.compile(r"Provided: (\d+), needed: (\d+)") - -# Insufficient-melt-inputs failures differ across mint implementations. 11005 is -# the registered "Transaction is not balanced" code (cdk), specific enough to -# trust on the code alone. 11000 is nutshell's generic, unregistered -# TransactionError covering many unrelated failures, so it only counts as a fee -# shortfall alongside the "not enough inputs" detail text. With no code suffix at -# all, that same text is the only signal. - - -def _net_minted_amount(amount_msat: int, token_unit: str, fees: int) -> int: - """ - Convert the token value minus fees (given in the token unit) into an - amount in the primary mint's unit. - """ - fee_msat = _sats_to_msats(fees) if token_unit == "sat" else fees - remaining_msat = amount_msat - fee_msat - if settings.primary_mint_unit == "sat": - return _msats_to_sats(remaining_msat) - return int(remaining_msat) - - -def _melt_definitively_failed(error: Exception) -> bool: - """Return whether the mint authoritatively rejected the Lightning payment. - - Cashu releases the reserved proofs for these responses, so the token remains - reusable. Transport failures and unknown errors are deliberately excluded: - after dispatch their payment outcome may still be pending or paid. - """ - message = str(error).strip() - return message.lower() == "could not pay invoice." or "(Code: 20004)" in message - - -def _melt_insufficient_shortfall(error: Exception) -> int | None: - """ - Classify a melt failure: return the observed shortfall (in the token unit) - when the mint rejected the inputs as insufficient, or None when the failure - is unrelated to fees and must not be retried (e.g. a Lightning payment - failure, where a smaller invoice would not help). - - Cashu errors carry no structured amounts (NUT-00 defines only detail/code, - flattened to "Mint Error: (Code: )" by cashu-py), so the - classification uses the code and the shortfall must be inferred: the - "Provided: X, needed: Y" amounts are nutshell-specific free text and only - refine the shortfall when present; otherwise shrink one unit at a time. - """ - message = str(error) - code_match = _MINT_ERROR_CODE_RE.search(message) - code = code_match.group(1) if code_match is not None else None - has_shortfall_text = "not enough inputs" in message.lower() - - match code: - case "11005": # registered TransactionUnbalanced: trust the code - pass - case "11000" if has_shortfall_text: # generic nutshell error: needs the text - pass - case None if has_shortfall_text: # no code suffix: text is the only signal - pass - case _: # other codes, a bare 11000, or no signal: must not retry - return None - - amounts = _MELT_SHORTFALL_RE.search(message) - if amounts is not None: - provided, needed = int(amounts.group(1)), int(amounts.group(2)) - if needed > provided: - return needed - provided - return 1 - - def _trusted_destination_candidates( candidates: list[str] | None = None, ) -> list[str]: @@ -983,692 +1010,6 @@ def _trusted_destination_candidates( return selected -async def _request_mint_with_fallback( - amount: int, - *, - op_name: str, - primary_wallet: Wallet | None = None, - destination_mints: list[str] | None = None, -) -> tuple[Wallet, str, MintQuote]: - """Try request_mint on the primary mint, fall back to other trusted mints - on transport or rate-limit failure. Returns the wallet, mint_url, and quote. - - Guards against amount <= 0: the cashu library's PostMintQuoteRequest - enforces ``amount > 0`` (Pydantic Field(gt=0)), so passing 0 raises a - cryptic validation error deep in the stack. Fail fast with context. - """ - if amount <= 0: - raise ValueError( - f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. " - f"Token value is too small after fee deduction or unit conversion." - ) - candidates = _trusted_destination_candidates(destination_mints) - logger.info( - "Trying trusted destination mints", - extra={ - "event": "cashu_destination_candidates", - "op_name": op_name, - "amount": amount, - "unit": settings.primary_mint_unit, - "candidates": candidates, - }, - ) - tried: list[str] = [] - for candidate_index, mint_url in enumerate(candidates, start=1): - cooldown = mint_cooldown_remaining(mint_url) - if cooldown > 0: - tried.append(f"{mint_url}: cooling down") - logger.warning( - "Skipping unavailable destination mint", - extra={ - "event": "cashu_destination_skipped", - "mint_url": mint_url, - "cooldown_seconds": round(cooldown, 2), - "op_name": op_name, - "candidate_index": candidate_index, - "candidate_count": len(candidates), - }, - ) - continue - logger.info( - "Trying destination mint", - extra={ - "event": "cashu_destination_attempt", - "mint_url": mint_url, - "op_name": op_name, - "candidate_index": candidate_index, - "candidate_count": len(candidates), - }, - ) - try: - if mint_url == settings.primary_mint and primary_wallet is not None: - wallet = primary_wallet - else: - wallet = await get_wallet( - mint_url, - settings.primary_mint_unit, - retry_on_rate_limit=False, - ) - quote = await run_mint_operation( - lambda: wallet.request_mint(amount), - op_name=op_name, - mint_url=mint_url, - retry_timeouts=False, - retry_on_rate_limit=False, - ) - logger.info( - "Destination mint selected", - extra={ - "event": "cashu_destination_selected", - "mint_url": mint_url, - "op_name": op_name, - "candidate_index": candidate_index, - "fallback_used": candidate_index > 1, - }, - ) - return wallet, mint_url, quote - except Exception as error: - tried.append(f"{mint_url}: {type(error).__name__}") - connection_failure = is_mint_connection_error(error) - rate_limited = is_mint_rate_limited(error) - if not connection_failure and not rate_limited: - raise - if connection_failure: - MintRateGuard.get(mint_url).apply_cooldown( - MINT_TRANSPORT_COOLDOWN_SECONDS, reason="unreachable" - ) - logger.warning( - "Destination mint failed", - extra={ - "event": "cashu_destination_failed", - "failed_mint": mint_url, - "error": str(error), - "error_type": type(error).__name__, - "connection_failure": connection_failure, - "rate_limited": rate_limited, - "tried": tried, - "op_name": op_name, - "candidate_index": candidate_index, - "candidate_count": len(candidates), - }, - ) - continue - logger.error( - "All trusted destination mints failed", - extra={ - "event": "cashu_destination_exhausted", - "op_name": op_name, - "amount": amount, - "unit": settings.primary_mint_unit, - "candidates": candidates, - "tried": tried, - }, - ) - raise MintConnectionError(f"All mints failed for {op_name}: {tried}") - - -async def _calculate_swap_amount( - amount_msat: int, - token_unit: str, - token_mint_url: str, - token_wallet: Wallet, - primary_wallet: Wallet | None, - proofs: list, - destination_mints: list[str] | None = None, -) -> int: - """ - Calculate the amount to mint on the primary mint after accounting for - melt fees and NUT-02 input fees on the foreign mint. - """ - if settings.primary_mint_unit == "sat": - receive_amount = _msats_to_sats(amount_msat) - else: - receive_amount = amount_msat - - if token_mint_url == settings.primary_mint: - logger.info( - "swap_to_trusted_mint: skipping fee estimation (same mint)", - extra={"minted_amount": receive_amount}, - ) - return int(receive_amount) - - # The cashu library's PostMintQuoteRequest enforces amount > 0 (Pydantic - # Field(gt=0)). When the token's face value in the primary mint's unit - # truncates to 0 (e.g. < 1000 msat with a "sat" primary unit), calling - # request_mint(0) raises a validation error that is cryptic in production - # logs. Guard early with full diagnostic context instead. - if receive_amount <= 0: - logger.error( - "swap_to_trusted_mint: receive_amount is zero or negative, cannot estimate fees", - extra={ - "amount_msat": amount_msat, - "token_unit": token_unit, - "token_mint_url": token_mint_url, - "primary_mint": settings.primary_mint, - "primary_mint_unit": settings.primary_mint_unit, - "receive_amount": receive_amount, - }, - ) - raise ValueError( - f"Token amount ({amount_msat} msat, unit={token_unit}) is too small to " - f"swap to primary mint ({settings.primary_mint}, unit={settings.primary_mint_unit}): " - f"receive_amount={receive_amount}. Minimum 1 {settings.primary_mint_unit} required." - ) - - logger.info( - "swap_to_trusted_mint: estimating fees", - extra={ - "dummy_amount": receive_amount, - "unit": settings.primary_mint_unit, - "token_mint_url": token_mint_url, - "primary_mint": settings.primary_mint, - "amount_msat": amount_msat, - }, - ) - - stage = "destination_fee_quote" - try: - _, _, dummy_mint_quote = await _request_mint_with_fallback( - receive_amount, - op_name="swap_fee_est_mint_quote", - primary_wallet=primary_wallet, - destination_mints=destination_mints, - ) - stage = "source_fee_quote" - dummy_melt_quote = await run_mint_operation( - lambda: token_wallet.melt_quote(dummy_mint_quote.request), - op_name="swap_fee_est_melt_quote", - mint_url=token_mint_url, - ) - - fee_reserve = dummy_melt_quote.fee_reserve - input_fees = token_wallet.get_fees_for_proofs(proofs) - total_fees = fee_reserve + input_fees - minted_amount = _net_minted_amount(amount_msat, token_unit, total_fees) - - if minted_amount <= 0: - raise ValueError(f"Fees ({total_fees} {token_unit}) exceed token amount") - - logger.info( - "swap_to_trusted_mint: fee estimation result", - extra={ - "token_amount_sat": _msats_to_sats(amount_msat), - "estimated_fee": total_fees, - "estimated_fee_unit": token_unit, - "input_fees": input_fees, - "minted_amount": minted_amount, - "minted_unit": settings.primary_mint_unit, - "fee_reserve": fee_reserve, - "token_mint_url": token_mint_url, - "primary_mint": settings.primary_mint, - }, - ) - return minted_amount - - except Exception as e: - logger.error( - "Cashu swap fee estimation failed", - extra={ - "event": "cashu_swap_fee_estimation_failed", - "stage": stage, - "error": str(e), - "error_type": type(e).__name__, - "amount_msat": amount_msat, - "token_unit": token_unit, - "token_mint_url": token_mint_url, - "primary_mint": settings.primary_mint, - "primary_mint_unit": settings.primary_mint_unit, - "receive_amount": receive_amount, - }, - ) - if is_mint_connection_error(e): - if stage == "source_fee_quote": - logger.error( - "Source mint is unreachable; destination fallback cannot spend its proofs", - extra={ - "event": "cashu_source_mint_unreachable", - "source_mint": token_mint_url, - "stage": stage, - "fallback_possible": False, - "reason": "cashu_proofs_are_bound_to_the_issuing_mint", - }, - ) - raise SourceMintConnectionError( - "Issuing Cashu mint is unreachable" - ) from e - raise MintConnectionError("Cashu mint is unreachable") from e - raise ValueError(f"Failed to estimate fees: {e}") from e - - -async def _reconcile_ambiguous_melt( - wallet: Wallet, quote_id: str, proofs: list[Proof] -) -> bool: - """Confirm a dispatched melt is paid or conservatively mark it ambiguous. - - A PAID quote is authoritative and does not require a proof-state lookup. - Every other immediate snapshot remains unsafe to retry: an in-flight - Lightning payment can still move UNPAID/UNSPENT to PENDING or PAID after the - cancelled HTTP request returns. - """ - try: - quote = await run_mint_operation( - lambda: wallet.get_melt_quote(quote_id), - op_name="reconcile_swap_melt_quote", - mint_url=str(wallet.url), - retry_timeouts=False, - ) - except Exception as error: - raise TokenConsumedError( - "Source melt outcome is unknown; reconciliation required" - ) from error - - if quote is not None and quote.state == MeltQuoteState.paid: - return True - - try: - proof_response = await run_mint_operation( - lambda: wallet.check_proof_state(proofs), - op_name="reconcile_swap_proofs", - mint_url=str(wallet.url), - retry_timeouts=False, - ) - proof_states = [state.state.value for state in proof_response.states] - except Exception: - proof_states = [] - - quote_state = getattr(getattr(quote, "state", None), "value", "unknown") - raise TokenConsumedError( - "Source melt outcome is ambiguous; reconciliation required " - f"(quote_state={quote_state}, proof_states={proof_states})" - ) - - -async def _confirm_melt_paid( - wallet: Wallet, quote_id: str, proofs: list[Proof], response: object -) -> bool: - """Accept a melt response only when PAID is explicit or reconciled.""" - if getattr(response, "state", None) == MeltQuoteState.paid: - return True - return await _reconcile_ambiguous_melt(wallet, quote_id, proofs) - - -async def swap_to_trusted_mint( - token_obj: Token, - token_wallet: Wallet, - *, - destination_mints: list[str] | None = None, -) -> tuple[int, str, str]: - logger.info( - "Starting Cashu cross-mint swap", - extra={ - "event": "cashu_swap_started", - "source_mint": token_obj.mint, - "token_amount": token_obj.amount, - "unit": token_obj.unit, - "primary_mint": settings.primary_mint, - }, - ) - # Ensure amount is an integer - if not isinstance(token_obj.amount, int): - token_amount = int(token_obj.amount) - else: - token_amount = token_obj.amount - - if token_obj.unit == "sat": - amount_msat = _sats_to_msats(token_amount) - elif token_obj.unit == "msat": - amount_msat = token_amount - else: - raise ValueError("Invalid unit") - destination_candidates = _trusted_destination_candidates(destination_mints) - # If the token is already from an allowed destination, redeem it same-mint. - # There's no melt/Lightning fee, but the mint's NUT-02 input fee still - # applies; _redeem_same_mint accounts for it. - if token_obj.mint in destination_candidates: - logger.info( - "swap_to_trusted_mint: token already on primary mint, skipping swap", - extra={ - "mint": token_obj.mint, - "amount": token_amount, - "unit": token_obj.unit, - }, - ) - return await _redeem_same_mint(token_wallet, token_obj) - - try: - proofs = await _load_and_resolve_token_proofs( - token_wallet, - token_obj, - op_name="swap_load_source_mint", - ) - except Exception as error: - if is_mint_connection_error(error): - raise SourceMintConnectionError( - "Issuing Cashu mint is unreachable" - ) from error - raise - - primary_wallet: Wallet | None = None - - minted_amount = await _calculate_swap_amount( - amount_msat, - token_obj.unit, - token_obj.mint, - token_wallet, - primary_wallet, - proofs, - destination_candidates, - ) - - # The estimate above is non-binding: the mint may demand a higher fee on the - # real quote or reject the melt outright. Retry the quote/melt cycle with the - # amount recomputed from the fees the mint actually demands. - observed_extra_fee = 0 - attempt = 0 - dest_wallet = primary_wallet - dest_mint_url = settings.primary_mint - while True: - attempt += 1 - if minted_amount <= 0: - logger.error( - "swap_to_trusted_mint: minted_amount is zero or negative before requesting quote", - extra={ - "minted_amount": minted_amount, - "attempt": attempt, - "foreign_mint": token_obj.mint, - "token_amount": token_amount, - "token_unit": token_obj.unit, - "amount_msat": amount_msat, - "observed_extra_fee": observed_extra_fee, - "primary_mint": settings.primary_mint, - }, - ) - raise ValueError( - f"Cannot swap token ({token_amount} {token_obj.unit}) from {token_obj.mint}: " - f"minted_amount={minted_amount} after fee deduction (attempt {attempt})" - ) - dest_wallet, dest_mint_url, mint_quote = await _request_mint_with_fallback( - minted_amount, - op_name="swap_request_mint", - primary_wallet=primary_wallet, - destination_mints=destination_candidates, - ) - logger.info( - "swap_to_trusted_mint: mint quote received", - extra={ - "mint_quote_id": mint_quote.quote, - "attempt": attempt, - "dest_mint": dest_mint_url, - }, - ) - - logger.info( - "Requesting melt quote from source mint", - extra={ - "event": "cashu_source_melt_quote_attempt", - "source_mint": token_obj.mint, - "destination_mint": dest_mint_url, - "attempt": attempt, - }, - ) - try: - melt_quote = await run_mint_operation( - lambda: token_wallet.melt_quote(mint_quote.request), - op_name="swap_melt_quote", - mint_url=token_obj.mint, - ) - except Exception as error: - if is_mint_connection_error(error): - logger.error( - "Source mint is unreachable; destination fallback cannot spend its proofs", - extra={ - "event": "cashu_source_mint_unreachable", - "source_mint": token_obj.mint, - "destination_mint": dest_mint_url, - "stage": "source_melt_quote", - "error": str(error), - "error_type": type(error).__name__, - "attempt": attempt, - }, - ) - raise SourceMintConnectionError( - "Issuing Cashu mint is unreachable" - ) from error - raise - input_fees = token_wallet.get_fees_for_proofs(proofs) - total_needed = melt_quote.amount + melt_quote.fee_reserve + input_fees - logger.info( - "swap_to_trusted_mint: melt quote received", - extra={ - "melt_quote_id": melt_quote.quote, - "melt_amount": melt_quote.amount, - "melt_fee_reserve": melt_quote.fee_reserve, - "input_fees": input_fees, - "total_needed": total_needed, - "token_amount": token_amount, - "attempt": attempt, - }, - ) - - if total_needed > token_amount: - recomputed = _net_minted_amount( - amount_msat, - token_obj.unit, - melt_quote.fee_reserve + input_fees + observed_extra_fee, - ) - if attempt >= _MAX_SWAP_ATTEMPTS or recomputed <= 0: - logger.warning( - "swap_to_trusted_mint: insufficient token amount for melt fees", - extra={ - "token_amount": token_amount, - "melt_amount": melt_quote.amount, - "melt_fee_reserve": melt_quote.fee_reserve, - "input_fees": input_fees, - "total_needed": total_needed, - "shortfall": total_needed - token_amount, - "attempts": attempt, - }, - ) - raise ValueError( - f"Token amount ({token_amount} {token_obj.unit}) is insufficient to cover " - f"melt fees. Needed: {total_needed} {token_obj.unit} " - f"(amount: {melt_quote.amount} + fee: {melt_quote.fee_reserve} + input_fees: {input_fees})" - ) - logger.warning( - "swap_to_trusted_mint: melt quote exceeds token amount, retrying", - extra={ - "total_needed": total_needed, - "token_amount": token_amount, - "retry_minted_amount": recomputed, - "attempt": attempt, - }, - ) - minted_amount = recomputed - continue - - try: - melt_response = await run_mint_operation( - lambda: token_wallet.melt( - proofs=proofs, - invoice=mint_quote.request, - fee_reserve_sat=melt_quote.fee_reserve, - quote_id=melt_quote.quote, - ), - op_name="swap_melt", - mint_url=token_obj.mint, - retry_timeouts=False, - ) - await _confirm_melt_paid( - token_wallet, melt_quote.quote, proofs, melt_response - ) - except Exception as e: - shortfall = _melt_insufficient_shortfall(e) - if shortfall is None: - if isinstance(e, TokenConsumedError): - raise - if _melt_definitively_failed(e): - raise ValueError( - f"Failed to melt token from foreign mint {token_obj.mint}: {e}" - ) from e - if is_mint_connection_error(e): - await _reconcile_ambiguous_melt( - token_wallet, melt_quote.quote, proofs - ) - logger.info( - "Source melt reconciled as paid; minting on destination", - extra={ - "event": "cashu_source_melt_reconciled_paid", - "source_mint": token_obj.mint, - "destination_mint": dest_mint_url, - "melt_quote_id": melt_quote.quote, - }, - ) - break - raise TokenConsumedError( - "Source melt failed after dispatch; outcome requires reconciliation" - ) from e - - observed_extra_fee += shortfall - recomputed = _net_minted_amount( - amount_msat, - token_obj.unit, - melt_quote.fee_reserve + input_fees + observed_extra_fee, - ) - if attempt >= _MAX_SWAP_ATTEMPTS or recomputed <= 0: - logger.error( - "swap_to_trusted_mint: melt failed", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "foreign_mint": token_obj.mint, - "token_amount": token_amount, - "melt_quote_id": melt_quote.quote, - "total_needed": total_needed, - "attempts": attempt, - }, - ) - raise ValueError( - f"Failed to melt token from foreign mint {token_obj.mint}: {e}" - ) from e - logger.warning( - "swap_to_trusted_mint: mint demanded more than quoted at melt, retrying", - extra={ - "shortfall": shortfall, - "retry_minted_amount": recomputed, - "attempt": attempt, - }, - ) - minted_amount = recomputed - continue - - break - - logger.info( - "Source melt succeeded; minting on destination", - extra={ - "event": "cashu_destination_mint_attempt", - "minted_amount": minted_amount, - "mint_quote_id": mint_quote.quote, - "dest_mint": dest_mint_url, - }, - ) - - await dest_wallet.load_proofs(reload=True) - pre_mint_balance = dest_wallet.available_balance.amount - try: - _ = await run_mint_operation( - lambda: dest_wallet.mint(minted_amount, quote_id=mint_quote.quote), - op_name="swap_mint_on_destination", - mint_url=dest_mint_url, - retry_timeouts=False, - ) - except Exception as e: - if "11003" in str(e) or "outputs already signed" in str(e).lower(): - # Previous mint call signed outputs at the mint but failed before - # bump_secret_derivation ran locally. Recover orphaned proofs and - # advance the counter so the next request derives fresh secrets. - logger.warning( - "swap_to_trusted_mint: outputs already signed — recovering orphaned proofs", - extra={ - "mint_quote_id": mint_quote.quote, - "minted_amount": minted_amount, - }, - ) - try: - for keyset_id in dest_wallet.keysets: - await dest_wallet.restore_tokens_for_keyset( - keyset_id, to=1, batch=25 - ) - await dest_wallet.load_proofs(reload=True) - post_recovery_balance = dest_wallet.available_balance.amount - balance_gained = post_recovery_balance - pre_mint_balance - logger.info( - "swap_to_trusted_mint: recovery scan completed", - extra={ - "pre_mint_balance": pre_mint_balance, - "post_recovery_balance": post_recovery_balance, - "balance_gained": balance_gained, - "expected": minted_amount, - }, - ) - if balance_gained < minted_amount: - # Recovery scan ran but did NOT restore the orphaned proofs - # (mint reports them as spent — they're stuck). Refuse to - # credit the API key balance for proofs we don't actually hold. - raise TokenConsumedError( - f"Swap recovery failed: mint signed outputs but proofs are " - f"unrecoverable (mint reports them spent). " - f"Expected {minted_amount}, recovered {balance_gained}. " - f"Local wallet DB ('.wallet/') state is corrupted — " - f"the counter for keyset is stuck at a bad index range." - ) - except TokenConsumedError: - raise - except Exception as recovery_err: - logger.error( - "swap_to_trusted_mint: recovery failed", - extra={"error": str(recovery_err)}, - ) - raise TokenConsumedError( - f"Mint on primary failed and recovery unsuccessful: {e}" - ) from e - else: - logger.error( - "swap_to_trusted_mint: mint on primary failed after successful melt", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "minted_amount": minted_amount, - "mint_quote_id": mint_quote.quote, - }, - ) - # Foreign proofs already melted (spent) — non-retryable. - raise TokenConsumedError( - "Mint on primary failed after successful melt" - ) from e - - logger.info( - "Cashu cross-mint swap completed", - extra={ - "event": "cashu_swap_completed", - "source_mint": token_obj.mint, - "dest_mint": dest_mint_url, - "original_amount": token_amount, - "minted_amount": minted_amount, - "unit": settings.primary_mint_unit, - }, - ) - - return int(minted_amount), settings.primary_mint_unit, dest_mint_url - - -async def swap_to_primary_mint( - token_obj: Token, token_wallet: Wallet -) -> tuple[int, str, str]: - """Backward-compatible alias for callers using the old function name.""" - return await swap_to_trusted_mint(token_obj, token_wallet) - - async def credit_balance( cashu_token: str, key: db.ApiKey, session: db.AsyncSession ) -> int: @@ -1688,10 +1029,8 @@ async def _credit_balance_locked( ) try: - destination_mint = key.refund_mint_url or settings.primary_mint amount, unit, mint_url = await recieve_token( cashu_token, - destination_mint=destination_mint, destination_unit=key.refund_currency if isinstance(key.refund_currency, str) else None, @@ -1788,21 +1127,35 @@ async def _credit_balance_locked( ) return amount except Exception as e: - logger.error( - "credit_balance: Error during token redemption", - extra={"error": str(e), "error_type": type(e).__name__}, + classification = classify_redemption_error(e) + expected_codes = { + "cashu_token_already_spent", + "cashu_source_mint_unreachable", + "cashu_mint_unreachable", + "cashu_mint_rate_limited", + "cashu_mint_timeout", + } + log = ( + logger.info + if classification is not None and classification[3] in expected_codes + else logger.error + ) + log( + "credit_balance: Token redemption failed", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "error_code": classification[3] if classification else None, + }, ) raise _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] = {} _wallet_load_locks: dict[str, asyncio.Lock] = {} -# Minimum seconds between full mint info + proof reloads for the same -# wallet. Prevents redundant mint API calls when get_wallet(load=True) -# is called rapidly by multiple background tasks (balance fetch, payout, -# auto-topup all hitting get_wallet within the same cycle). -_WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30 async def get_wallet( @@ -1811,8 +1164,9 @@ async def get_wallet( load: bool = True, retry_on_rate_limit: bool = True, force_reload: bool = False, + load_proofs: bool = True, ) -> Wallet: - global _wallets, _wallet_last_load, _wallet_load_locks + global _wallets, _wallet_last_load, _wallet_last_mint_load, _wallet_load_locks id = f"{mint_url}_{unit}" lock = _wallet_load_locks.setdefault(id, asyncio.Lock()) async with lock: @@ -1821,25 +1175,39 @@ async def get_wallet( if load: now = time.monotonic() - last = _wallet_last_load.get(id) + last_mint_load = _wallet_last_mint_load.get(id) if ( force_reload - or last is None - or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS + 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(), + 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, ) - await run_mint_operation( - lambda: _wallets[id].load_proofs(reload=True), - op_name="load_proofs", - mint_url=mint_url, - retry_on_rate_limit=retry_on_rate_limit, - ) - _wallet_last_load[id] = time.monotonic() + _wallet_last_mint_load[id] = time.monotonic() + + if load_proofs: + last_proof_load = _wallet_last_load.get(id) + if ( + force_reload + or last_proof_load is None + or now - last_proof_load + >= _WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS + ): + await run_mint_operation( + lambda: _wallets[id].load_proofs(reload=True), + op_name="load_proofs", + mint_url=mint_url, + retry_on_rate_limit=retry_on_rate_limit, + ) + _wallet_last_load[id] = time.monotonic() return _wallets[id] @@ -1912,13 +1280,14 @@ async def _get_supported_mint_units(mint_url: str) -> list[str]: if cached is not None and now < cached[0]: return cached[1] - wallet = await get_wallet(mint_url, settings.primary_mint_unit, load=False) - keysets = await run_mint_operation( - lambda: wallet._get_keysets(), - op_name="get_mint_keysets", - mint_url=mint_url, + # A metadata load populates Cashu's shared keyset cache for all units. + wallet = await get_wallet( + mint_url, + settings.primary_mint_unit, retry_on_rate_limit=False, + load_proofs=False, ) + keysets = await get_cashu_keysets(mint_url=wallet.url, db=wallet.db) units: list[str] = [] for keyset in keysets: if not keyset.active or keyset.unit is None: @@ -2568,11 +1937,18 @@ async def periodic_routstr_fee_payout() -> None: logger.warning("Routstr fee payout was already claimed") continue except BaseException as e: - logger.critical( - "Routstr fee payout outcome is unknown; awaiting quote reconciliation", - extra={"payout_in_progress_msats": paid_msats}, - exc_info=isinstance(e, Exception), - ) + if attempt_quote_id is None: + logger.error( + "Routstr fee payout failed before melt dispatch", + extra={"payout_msats": paid_msats}, + exc_info=isinstance(e, Exception), + ) + else: + logger.critical( + "Routstr fee payout outcome is unknown; awaiting quote reconciliation", + extra={"payout_in_progress_msats": paid_msats}, + exc_info=isinstance(e, Exception), + ) if not isinstance(e, Exception): raise continue @@ -2617,13 +1993,42 @@ async def periodic_routstr_fee_payout() -> None: ) -async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: +def _quote_callback( + notify: Callable[[str, str], Awaitable[None]], mint: str +) -> Callable[[str], Awaitable[None]]: + async def callback(quote_id: str) -> None: + await notify(quote_id, mint) + + return callback + + +async def send_to_lnurl( + amount: int, + unit: str, + mint: str, + address: str, + *, + on_melt_quote: Callable[[str, str], Awaitable[None]] | None = None, +) -> int: + """``on_melt_quote`` gets the quote id and the mint that issued it, since + fallback may pick a different mint than requested.""" async with wallet_operation_guard(): mint = await find_trusted_mint_with_funds(amount, unit, mint, force_reload=True) wallet = await get_wallet(mint, unit) available = get_proofs_per_mint_and_unit(wallet, mint, unit, not_reserved=True) - proofs, _ = await wallet.select_to_send(available, amount, set_reserved=True) - return await raw_send_to_lnurl(wallet, proofs, address, unit) + # Hand over unreserved proofs: raw_send_to_lnurl reserves only once the + # destination, the invoice amount and the melt quote have all been + # accepted, so a rejected refund cannot strand locked proofs. + return await raw_send_to_lnurl( + wallet, + available, + address, + unit, + amount=amount, + on_melt_quote=( + None if on_melt_quote is None else _quote_callback(on_melt_quote, mint) + ), + ) # class Payment: diff --git a/scripts/refund_token_to_lightning.py b/scripts/refund_token_to_lightning.py new file mode 100644 index 00000000..7a7d0dae --- /dev/null +++ b/scripts/refund_token_to_lightning.py @@ -0,0 +1,89 @@ +"""Redeem a cashu token into a balance and pay it out to a Lightning address. + +Usage: + python scripts/refund_token_to_lightning.py [--url http://localhost:8000] + +Steps: + 1. POST /v1/balance/create redeems the token into a fresh API key + 2. POST /v1/balance/refund pays the full balance to the Lightning address + +A 502 from the refund means the melt was dispatched but unconfirmed; the +balance is withheld until the server reconciles it. Re-run with the printed +API key to check whether it settled. +""" + +import argparse +import ipaddress +import sys +from urllib.parse import urlparse + +import httpx + + +def _is_loopback(host: str) -> bool: + if host == "localhost": + return True + try: + return ipaddress.ip_address(host.strip("[]")).is_loopback + except ValueError: + return False + + +def check_url(url: str) -> str: + """Reject a URL that would put the token and the API key on the wire.""" + parsed = urlparse(url) + if parsed.scheme == "https": + return url + if parsed.scheme == "http" and _is_loopback(parsed.hostname or ""): + return url + raise SystemExit( + f"Refusing to send a cashu token and bearer key to {url!r}: " + "use https, or http only for a loopback host." + ) + + +def create_balance(client: httpx.Client, token: str) -> str: + response = client.post("/v1/balance/create", json={"initial_balance_token": token}) + response.raise_for_status() + data = response.json() + print(f"Redeemed token: balance {data['balance']} msats, key {data['api_key']}") + return str(data["api_key"]) + + +def refund_to_lightning(client: httpx.Client, api_key: str, address: str) -> dict: + response = client.post( + "/v1/balance/refund", + headers={"Authorization": f"Bearer {api_key}"}, + json={"lightning_address": address}, + ) + if response.status_code >= 400: + print(f"Refund failed ({response.status_code}): {response.text}") + sys.exit(1) + return dict(response.json()) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("token", help="cashu token, or sk-... key from a prior run") + parser.add_argument("lightning_address", help="Lightning address or LNURL") + parser.add_argument("--url", default="http://localhost:8000", help="routstr URL") + args = parser.parse_args() + + with httpx.Client(base_url=check_url(args.url), timeout=120.0) as client: + api_key = ( + args.token + if args.token.startswith("sk-") + else create_balance(client, args.token) + ) + result = refund_to_lightning(client, api_key, args.lightning_address) + + amount = result.get("sats") or result.get("msats") + unit = "sats" if "sats" in result else "msats" + print( + f"Refund {result['refund_id']} {result['status']}: " + f"{amount} {unit} -> {result.get('recipient', args.lightning_address)}" + ) + + +if __name__ == "__main__": + main() diff --git a/tests/conftest.py b/tests/conftest.py index 8112487d..d1bfa919 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -7,9 +7,27 @@ absent one) override this per-test via ``monkeypatch``. """ import os +from typing import Iterator + +import pytest # Valid Fernet keys; KEY_A is the suite default, KEY_B is for wrong-key tests. TEST_SECRET_KEY = "l_Tkp-7xmjcQ-IFhr6qhILrU8HPRbEmYMrfSbo_5srU=" TEST_SECRET_KEY_ALT = "_Teyrky_iToeDK51Tj1FsI9MJ340_cqKGmeher-a7MQ=" os.environ.setdefault("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY) + + +@pytest.fixture(autouse=True) +def _isolate_redemption_negative_cache() -> Iterator[None]: + """Clear the process-wide negative cache between tests. + + The cache deliberately persists terminal redemption failures across + requests; without this fixture a test that burns a token would poison + every later test reusing the same token string. + """ + from routstr.redemption_cache import redemption_negative_cache + + redemption_negative_cache.clear() + yield + redemption_negative_cache.clear() diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index aa10a81c..f8babb17 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -1,7 +1,7 @@ import asyncio import json import os -from typing import Any, AsyncGenerator, Callable, Dict, List, Optional, Tuple +from typing import Any, AsyncGenerator, Callable, Dict, Iterator, List, Optional, Tuple from unittest.mock import MagicMock, patch import pytest @@ -68,6 +68,14 @@ os.environ.pop("ADMIN_PASSWORD", None) from routstr.core.db import ApiKey, get_session # noqa: E402 from routstr.core.main import app, lifespan # noqa: E402 +from routstr.mint import MintRateGuard # noqa: E402 + + +@pytest.fixture(autouse=True) +def isolate_mint_rate_guards() -> Iterator[None]: + MintRateGuard._guards.clear() + yield + MintRateGuard._guards.clear() @pytest.fixture(scope="session") @@ -508,6 +516,10 @@ async def integration_app( # Copy all routes from the main app test_app.router = app.router + # ...and its exception handlers, so a request that fails here fails the way + # it would in production rather than escaping as a bare exception. + test_app.exception_handlers.update(app.exception_handlers) + # Override the get_session dependency async def override_get_session() -> AsyncGenerator[AsyncSession, None]: yield integration_session @@ -542,8 +554,8 @@ async def integration_app( patch("routstr.wallet.send_to_lnurl", testmint_wallet.send_to_lnurl), patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token), patch("routstr.wallet.get_balance", testmint_wallet.get_balance), - patch("routstr.balance.send_token", testmint_wallet.send_token), - patch("routstr.balance.send_to_lnurl", testmint_wallet.send_to_lnurl), + patch("routstr.refund.send_token", testmint_wallet.send_token), + patch("routstr.refund.send_to_lnurl", testmint_wallet.send_to_lnurl), patch("websockets.connect") as mock_websockets, patch("routstr.payment.price.btc_usd_price", return_value=50000.0), patch("routstr.payment.price.sats_usd_price", return_value=0.0005), diff --git a/tests/integration/test_admin_pricing_rate_validation.py b/tests/integration/test_admin_pricing_rate_validation.py new file mode 100644 index 00000000..fb6cfa18 --- /dev/null +++ b/tests/integration/test_admin_pricing_rate_validation.py @@ -0,0 +1,476 @@ +"""Admin write edge: a rate that is not a number never becomes a stored price. + +A billable rate is usable only when it is finite and non-negative. The admin +model endpoints are an entry point for rates the node will later bill on, and +they accept whatever a client sends: ``json`` parses the bare ``NaN``/ +``Infinity`` literals into real floats and overflows ``1e999`` to ``inf``, a +non-numeric string coerced silently to ``$0``, and a negative rate is truthy so +it read back as a chargeable price that bills a negative amount. + +These tests assert the edge answers a malformed rate with a 422 — a client bug +reported as a client bug — rather than persisting it or failing as a 500, and +that the operator can still open the listing that shows the row needing repair. +""" + +from __future__ import annotations + +import json +from datetime import datetime, timedelta, timezone + +import pytest +from httpx import AsyncClient +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import admin_sessions +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.proxy import reinitialize_upstreams + + +def _admin_headers() -> dict[str, str]: + token = "test-admin-rate-validation-token" + admin_sessions[token] = int( + (datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp() + ) + return {"Authorization": f"Bearer {token}"} + + +def _pricing(**overrides: object) -> dict[str, object]: + pricing: dict[str, object] = { + "prompt": 1.4e-7, + "completion": 2.8e-7, + "request": 0.0, + "image": 0.0, + "web_search": 0.0, + "internal_reasoning": 0.0, + "input_cache_read": 0.0, + "input_cache_write": 0.0, + } + pricing.update(overrides) + return pricing + + +def _payload( + provider_id: int, + *, + model_id: str = "rate-model", + pricing: dict[str, object] | None = None, +) -> dict[str, object]: + return { + "id": model_id, + "name": "Rate Model", + "description": "d", + "created": 0, + "context_length": 128000, + "architecture": { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + }, + "pricing": pricing if pricing is not None else _pricing(), + "per_request_limits": None, + "top_provider": None, + "upstream_provider_id": provider_id, + "canonical_slug": None, + "alias_ids": [], + "enabled": True, + "forwarded_model_id": model_id, + } + + +async def _make_provider(session: AsyncSession) -> int: + provider = UpstreamProviderRow( + provider_type="generic", + base_url="https://rate-upstream.example/v1", + api_key="test-key", + provider_fee=1.0, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + await reinitialize_upstreams() + assert provider.id is not None + return provider.id + + +def _raw_model_body(provider_id: int, model_id: str, prompt_literal: str) -> str: + """A request body built as text, so it can carry a literal ``json`` accepts + but Python's own encoder would refuse to produce.""" + return ( + f'{{"id": "{model_id}", "name": "raw", "description": "d", "created": 0,' + ' "context_length": 8192, "architecture": {"modality": "text"},' + f' "pricing": {{"prompt": {prompt_literal}, "completion": 2.8e-7}},' + f' "upstream_provider_id": {provider_id}}}' + ) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_negative_price_is_rejected( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A negative rate is not a valid price — accepting it would persist a row + that bills a negative amount, which settlement subtracts from the balance. + Being truthy, it also reads back as a chargeable price. Reject at the edge + rather than silently storing it.""" + provider_id = await _make_provider(integration_session) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + json=_payload(provider_id, model_id="neg-price", pricing=_pricing(prompt=-1.0)), + ) + + assert resp.status_code == 422 + assert await integration_session.get(ModelRow, ("neg-price", provider_id)) is None + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_malformed_price_string_is_rejected( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A present non-numeric rate is a client bug: it coerces to ``$0`` on the + read path, producing an unpriced-looking row indistinguishable from a + deliberate free price. Surface it as a 422 instead of accepting it.""" + provider_id = await _make_provider(integration_session) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + json=_payload( + provider_id, model_id="bad-price", pricing=_pricing(prompt="oops") + ), + ) + + assert resp.status_code == 422 + assert await integration_session.get(ModelRow, ("bad-price", provider_id)) is None + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_boolean_price_is_rejected( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A JSON ``true`` coerces to a finite, positive ``1.0`` — a dollar per + token — so it passes every numeric guard. The write edge asks the same + coercion the catalog readers do, and answers a 422.""" + provider_id = await _make_provider(integration_session) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + json=_payload( + provider_id, model_id="bool-price", pricing=_pricing(prompt=True) + ), + ) + + assert resp.status_code == 422 + assert await integration_session.get(ModelRow, ("bool-price", provider_id)) is None + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model_id", "pricing"), + [ + ("null-prompt", _pricing(prompt=None)), + ("null-aux-rate", _pricing(image=None)), + ("no-prompt", {k: v for k, v in _pricing().items() if k != "prompt"}), + ], + ids=["null-required", "null-auxiliary", "absent-required"], +) +async def test_a_rate_that_is_not_there_is_rejected( + model_id: str, + pricing: dict[str, object], + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """A rate given as ``null``, or a required rate left out, is not a price. + + ``dict.get`` cannot tell the two apart and skipped both, so a row + ``Pricing`` cannot parse was committed and the response that reads it back + raised. + """ + provider_id = await _make_provider(integration_session) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + json=_payload(provider_id, model_id=model_id, pricing=pricing), + ) + + assert resp.status_code == 422 + assert await integration_session.get(ModelRow, (model_id, provider_id)) is None + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_an_absent_auxiliary_rate_is_still_accepted( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """Only ``prompt`` and ``completion`` are required; the rest carry defaults, + and a payload that omits them must still be accepted.""" + provider_id = await _make_provider(integration_session) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + json=_payload( + provider_id, + model_id="lean-price", + pricing={"prompt": 1.4e-7, "completion": 2.8e-7}, + ), + ) + + assert resp.status_code == 200 + assert await integration_session.get(ModelRow, ("lean-price", provider_id)) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_numeric_string_price_is_still_accepted( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """The stored pricing JSON has always accepted numeric strings, and the UI + round-trips rates through text fields. Rejecting a *malformed* rate must not + also reject a well-formed one that arrives spelled as a string.""" + provider_id = await _make_provider(integration_session) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + json=_payload( + provider_id, model_id="string-price", pricing=_pricing(prompt="0.000005") + ), + ) + + assert resp.status_code == 200 + row = await integration_session.get(ModelRow, ("string-price", provider_id)) + assert row is not None + assert json.loads(row.pricing)["prompt"] == "0.000005" + + +def test_non_finite_price_is_rejected_by_the_write_model() -> None: + """``NaN``/``±inf`` are not billable rates: the carrier every write endpoint + shares must reject them before they can be persisted and read back as a + chargeable price.""" + from pydantic import ValidationError + + from routstr.core.admin import ModelCreate + + for bad in (float("nan"), float("inf"), float("-inf")): + with pytest.raises(ValidationError): + ModelCreate.model_validate( + _payload(1, model_id="nonfinite", pricing=_pricing(prompt=bad)) + ) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_oversized_integer_price_is_rejected( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A JSON integer too large for a float is a client bug, not a server fault. + + ``float()`` raises ``OverflowError`` for it, and pydantic converts only + ``ValueError``/``AssertionError`` into validation errors, so it escaped the + edge as a 500. It must be answered with the same 422 as every other + unusable rate. + """ + provider_id = await _make_provider(integration_session) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers={**_admin_headers(), "Content-Type": "application/json"}, + content=_raw_model_body(provider_id, "huge-price", "9" * 400), + ) + + assert resp.status_code == 422 + assert await integration_session.get(ModelRow, ("huge-price", provider_id)) is None + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_non_finite_literal_price_is_rejected( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A bare ``Infinity``/``NaN`` literal gets the same 422 as any other rate. + + ``json`` accepts both literals, so the edge sees a real float and rejects + it — but pydantic echoes the offending value back in the error's ``input`` + field, and the response encoder runs with ``allow_nan=False``. Serializing + that reply raised "Out of range float values are not JSON compliant", so the + 422 escaped as a 500 and reported a client's bad rate as a server fault. + """ + provider_id = await _make_provider(integration_session) + + for literal in ("Infinity", "-Infinity", "NaN"): + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers={**_admin_headers(), "Content-Type": "application/json"}, + content=_raw_model_body(provider_id, "odd-price", literal), + ) + + assert resp.status_code == 422, literal + assert ( + await integration_session.get(ModelRow, ("odd-price", provider_id)) is None + ) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_non_finite_literal_price_is_rejected_in_batch_override( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """The batch path shares the same carrier, so it must answer 422 too.""" + provider_id = await _make_provider(integration_session) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/batch-override", + headers={**_admin_headers(), "Content-Type": "application/json"}, + content=( + '{"models": [' + + _raw_model_body(provider_id, "odd-batch", "Infinity") + + "]}" + ), + ) + + assert resp.status_code == 422 + assert await integration_session.get(ModelRow, ("odd-batch", provider_id)) is None + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_model_listing_shows_a_non_finite_stored_rate( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """The operator must be able to see the rate that needs fixing. + + The admin listing is the one view that still carries a row the served + catalog holds back, and its encoder rendered a stored ``Infinity`` as + ``null`` — indistinguishable from a rate the row never carried. + """ + provider_id = await _make_provider(integration_session) + integration_session.add( + ModelRow( + id="inf-rate", + name="inf-rate", + description="d", + created=0, + context_length=8192, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + } + ), + pricing=json.dumps({"prompt": float("inf"), "completion": 2e-06}), + upstream_provider_id=provider_id, + enabled=True, + forwarded_model_id="inf-rate", + ) + ) + await integration_session.commit() + + resp = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + ) + + assert resp.status_code == 200 + listed = {m["id"]: m for m in resp.json()["db_models"]} + assert listed["inf-rate"]["pricing"]["prompt"] == "inf" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_malformed_auxiliary_rate_is_rejected( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """Validation spans every billable rate, not just the token rates. + + ``prompt``/``completion`` are the rates most prices are built from, but the + request, image, search, reasoning and cache rates are billed too. A negative + or non-finite value in any of them is the same defect and must be answered + the same way. + """ + provider_id = await _make_provider(integration_session) + + for field, bad in ( + ("request", -1.0), + ("image", -0.5), + ("web_search", float("inf")), + ("internal_reasoning", float("nan")), + ("input_cache_read", -1e-06), + ("input_cache_write", float("-inf")), + ("completion", -1.0), + ): + # Send raw bytes rather than `json=`: httpx>=0.28 refuses to encode + # non-finite floats itself (allow_nan=False), but the point of this + # test is that the SERVER answers the bare NaN/Infinity literals + # with a 422, so the literals must still reach it. + body = json.dumps( + _payload( + provider_id, model_id="aux-rate", pricing=_pricing(**{field: bad}) + ) + ).encode("utf-8") + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers={**_admin_headers(), "content-type": "application/json"}, + content=body, + ) + + assert resp.status_code == 422, field + assert ( + await integration_session.get(ModelRow, ("aux-rate", provider_id)) is None + ), field + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_single_model_shows_a_non_finite_stored_rate( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """The single-model view answers like the listing it is opened from. + + It is the other view of a row the served-catalog backstop holds back, so it + has the same duty to name the rate that needs fixing rather than rendering + it as ``null``. + """ + provider_id = await _make_provider(integration_session) + integration_session.add( + ModelRow( + id="inf-one", + name="inf-one", + description="d", + created=0, + context_length=8192, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + } + ), + pricing=json.dumps({"prompt": float("inf"), "completion": 2e-06}), + upstream_provider_id=provider_id, + enabled=True, + forwarded_model_id="inf-one", + ) + ) + await integration_session.commit() + + resp = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/models/inf-one", + headers=_admin_headers(), + ) + + assert resp.status_code == 200 + assert resp.json()["pricing"]["prompt"] == "inf" diff --git a/tests/integration/test_child_keys.py b/tests/integration/test_child_keys.py deleted file mode 100644 index 5226d357..00000000 --- a/tests/integration/test_child_keys.py +++ /dev/null @@ -1,190 +0,0 @@ -import asyncio -import secrets -from typing import Any - -import pytest -from fastapi import HTTPException -from sqlmodel.ext.asyncio.session import AsyncSession - -from routstr.auth import adjust_payment_for_tokens, pay_for_request -from routstr.balance import ChildKeyRequest, create_child_key -from routstr.core.db import ApiKey, create_session -from routstr.core.settings import settings - - -@pytest.mark.asyncio -async def test_child_key_flow(integration_session: AsyncSession) -> None: - # 1. Create a parent key with balance - parent_raw = "parent_test_key_" + secrets.token_hex(4) - parent_key = ApiKey( - hashed_key=parent_raw, - balance=10000, # 10 sats - ) - integration_session.add(parent_key) - await integration_session.commit() - await integration_session.refresh(parent_key) - - # Mock settings - settings.child_key_cost = 1000 # 1 sat - - # 2. Call create_child_key - result = await create_child_key( - ChildKeyRequest(count=1), parent_key, integration_session - ) - - assert "api_keys" in result - assert result["cost_msats"] == 1000 - assert result["parent_balance"] == 9000 - - child_key_raw = result["api_keys"][0][3:] # remove sk- - - # 3. Verify child key exists in DB - child_key_db = await integration_session.get(ApiKey, child_key_raw) - assert child_key_db is not None - assert child_key_db.parent_key_hash == parent_key.hashed_key - assert child_key_db.balance == 0 - - # 4. Test payment with child key - cost = 500 - await pay_for_request(child_key_db, cost, integration_session) - - # Refresh keys - await integration_session.refresh(parent_key) - await integration_session.refresh(child_key_db) - - # Parent should be charged - assert parent_key.reserved_balance == 500 - assert parent_key.total_requests == 1 - - # Child should have total_requests incremented - assert child_key_db.total_requests == 1 - - # 5. Test adjustment - response_data = {"model": "test-model", "usage": {"total_tokens": 10}} - - # Mock calculate_cost - import routstr.auth - from routstr.payment.cost_calculation import CostData - - async def mock_calculate_cost(*args: Any, **kwargs: Any) -> CostData: - return CostData( - base_msats=0, input_msats=200, output_msats=200, total_msats=400 - ) - - # Patch calculate_cost - original_calculate_cost = routstr.auth.calculate_cost - routstr.auth.calculate_cost = mock_calculate_cost - - try: - adjustment = await adjust_payment_for_tokens( - child_key_db, response_data, integration_session, 500, None, None - ) - assert adjustment["total_msats"] == 400 - - # Refresh keys - await integration_session.refresh(parent_key) - await integration_session.refresh(child_key_db) - - # Parent should have updated balance and total_spent - assert parent_key.reserved_balance == 0 - assert parent_key.balance == 9000 - 400 - assert ( - parent_key.total_spent == 1400 - ) # 1000 for child key creation + 400 for request - - # Child should also have total_spent updated - assert child_key_db.total_spent == 400 - - finally: - routstr.auth.calculate_cost = original_calculate_cost - - -@pytest.mark.asyncio -async def test_child_key_insufficient_balance( - integration_session: AsyncSession, -) -> None: - parent_key = ApiKey( - hashed_key="poor_parent_" + secrets.token_hex(4), - balance=500, - ) - integration_session.add(parent_key) - await integration_session.commit() - await integration_session.refresh(parent_key) - - settings.child_key_cost = 1000 - - with pytest.raises(HTTPException) as exc: - await create_child_key( - ChildKeyRequest(count=1), parent_key, integration_session - ) - assert exc.value.status_code == 402 - - -@pytest.mark.asyncio -async def test_concurrent_child_key_creation_is_atomic( - patched_db_engine: None, -) -> None: - """Two concurrent create_child_key() calls with balance for exactly one must - result in exactly one success and one 402, with the parent balance deducted - only once.""" - child_key_cost = 1000 - settings.child_key_cost = child_key_cost - - parent_hash = f"parent_concurrent_{secrets.token_hex(8)}" - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=child_key_cost) - session.add(parent) - await session.commit() - - results: list[str] = [] - - async def attempt() -> None: - async with create_session() as session: - fresh_parent = await session.get(ApiKey, parent_hash) - assert fresh_parent is not None - try: - await create_child_key(ChildKeyRequest(count=1), fresh_parent, session) - results.append("success") - except HTTPException as exc: - assert exc.status_code == 402 - results.append("blocked") - - await asyncio.gather(attempt(), attempt()) - - assert sorted(results) == ["blocked", "success"], ( - f"Expected exactly one success and one 402, got: {results}" - ) - - async with create_session() as session: - final = await session.get(ApiKey, parent_hash) - assert final is not None - - assert final.balance == 0, ( - f"Balance should be fully deducted once: expected 0, got {final.balance}" - ) - assert final.total_spent == child_key_cost, ( - f"total_spent should equal one deduction: expected {child_key_cost}, " - f"got {final.total_spent}" - ) - - -@pytest.mark.asyncio -async def test_child_key_cannot_create_child(integration_session: AsyncSession) -> None: - parent_key = ApiKey( - hashed_key="parent_" + secrets.token_hex(4), - balance=10000, - ) - child_key = ApiKey( - hashed_key="child_" + secrets.token_hex(4), - balance=0, - parent_key_hash=parent_key.hashed_key, - ) - integration_session.add(parent_key) - integration_session.add(child_key) - await integration_session.commit() - await integration_session.refresh(child_key) - - with pytest.raises(HTTPException) as exc: - await create_child_key(ChildKeyRequest(count=1), child_key, integration_session) - assert exc.value.status_code == 400 - assert "Cannot create a child key for another child key" in str(exc.value.detail) diff --git a/tests/integration/test_child_keys_api.py b/tests/integration/test_child_keys_api.py deleted file mode 100644 index 1adb9496..00000000 --- a/tests/integration/test_child_keys_api.py +++ /dev/null @@ -1,102 +0,0 @@ -from typing import Any - -import pytest -from httpx import AsyncClient - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_wallet_info_returns_child_keys( - integration_client: AsyncClient, - authenticated_client: AsyncClient, - integration_session: Any, -) -> None: - """Test that GET /v1/wallet/info returns child keys for a parent key""" - - # 1. Get parent info to find its hashed_key - response = await authenticated_client.get("/v1/wallet/info") - assert response.status_code == 200 - parent_data = response.json() - parent_data["api_key"] - - # 2. Create child keys for this parent - # We need to use the parent's authentication for this - child_payload = {"count": 2, "balance_limit": 1000, "balance_limit_reset": "daily"} - create_response = await authenticated_client.post( - "/v1/wallet/child-key", json=child_payload - ) - assert create_response.status_code == 200 - create_data = create_response.json() - child_keys = create_data["api_keys"] - assert len(child_keys) == 2 - - # 3. Call /info again and check for child_keys - info_response = await authenticated_client.get("/v1/wallet/info") - assert info_response.status_code == 200 - info_data = info_response.json() - - assert "child_keys" in info_data - assert len(info_data["child_keys"]) == 2 - - # Verify child key details - for ck in info_data["child_keys"]: - assert ck["api_key"] in child_keys - assert ck["balance_limit"] == 1000 - assert ck["balance_limit_reset"] == "daily" - assert "total_spent" in ck - assert "total_requests" in ck - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_wallet_info_child_key_no_child_keys( - integration_client: AsyncClient, - authenticated_client: AsyncClient, - integration_session: Any, -) -> None: - """Test that GET /v1/wallet/info for a child key does NOT return child_keys""" - - # 1. Create a child key - child_payload = {"count": 1} - create_response = await authenticated_client.post( - "/v1/wallet/child-key", json=child_payload - ) - assert create_response.status_code == 200 - child_key = create_response.json()["api_keys"][0] - - # 2. Use the child key to get its info - integration_client.headers["Authorization"] = f"Bearer {child_key}" - info_response = await integration_client.get("/v1/wallet/info") - assert info_response.status_code == 200 - info_data = info_response.json() - parent_key = authenticated_client._test_api_key # type: ignore[attr-defined] - parent_key_hash = parent_key.removeprefix("sk-") - - assert info_data["is_child"] is True - assert "child_keys" not in info_data - assert "parent_key" not in info_data - assert info_data["parent_key_preview"] == parent_key_hash[:8] + "..." - assert info_data["parent_key_preview"] not in {parent_key, parent_key_hash} - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_account_info_root_returns_child_keys( - authenticated_client: AsyncClient, -) -> None: - """Test that GET / returns child keys for a parent key (root endpoint)""" - - # 1. Create a child key - child_payload = {"count": 1} - await authenticated_client.post("/v1/wallet/child-key", json=child_payload) - - # 2. Call root endpoint /v1/balance/ - # Note: routstr/balance.py defines router = APIRouter() - # and it is included in balance_router with prefix /v1/balance - # The endpoint is @router.get("/") - response = await authenticated_client.get("/v1/balance/") - assert response.status_code == 200 - data = response.json() - - assert "child_keys" in data - assert len(data["child_keys"]) >= 1 diff --git a/tests/integration/test_error_handling_edge_cases.py b/tests/integration/test_error_handling_edge_cases.py index 4587fd69..cfd8c221 100644 --- a/tests/integration/test_error_handling_edge_cases.py +++ b/tests/integration/test_error_handling_edge_cases.py @@ -31,7 +31,7 @@ class TestNetworkFailureScenarios: AsyncMock(side_effect=ConnectError("Mint service unavailable")), ), patch( - "routstr.balance.send_token", + "routstr.refund.send_token", AsyncMock(side_effect=ConnectError("Mint service unavailable")), ), ): diff --git a/tests/integration/test_failover_billing.py b/tests/integration/test_failover_billing.py index a0c31401..d67dc32b 100644 --- a/tests/integration/test_failover_billing.py +++ b/tests/integration/test_failover_billing.py @@ -17,7 +17,7 @@ from httpx import AsyncClient from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession -from routstr.core.db import ApiKey, ReservationRelease +from routstr.core.db import ReservationRelease from routstr.payment.models import Architecture, Model, Pricing from routstr.proxy import refresh_model_maps from routstr.upstream.base import BaseUpstreamProvider @@ -531,103 +531,6 @@ async def test_failover_beyond_balance_envelope_is_rejected( # fallback must be rejected before its upstream is ever contacted. assert response.status_code == 402 assert [r.url.host for r in sent_requests] == ["cheap.example.com"] - - -@pytest.fixture -async def three_candidate_child_maps( - patched_db_engine: None, -) -> AsyncGenerator[None, None]: - """Second candidate cannot fit the child limit; third restores and serves.""" - first = _StaticProvider( - CHEAP_BASE_URL, - "key-first", - 1.0, - _make_model("dual-model", 0.001, 0.002, max_cost=50.0), - ) - too_large = _StaticProvider( - EXPENSIVE_BASE_URL, - "key-too-large", - 1.0, - _make_model("dual-model", 0.002, 0.003, max_cost=100.0), - ) - third = _StaticProvider( - THIRD_BASE_URL, - "key-third", - 1.0, - _make_model("dual-model", 0.003, 0.004, max_cost=50.0), - ) - async for _ in _install_providers([first, too_large, third]): - yield - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_child_failover_rolls_back_failed_larger_reserve_before_restoring( - authenticated_client: AsyncClient, - three_candidate_child_maps: None, - integration_session: AsyncSession, -) -> None: - """A failed child guard cannot leak its parent update into restoration.""" - key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined] - child = await integration_session.get(ApiKey, key_hash) - assert child is not None - parent = ApiKey(hashed_key="failover-parent", balance=10_000_000) - child.parent_key_hash = parent.hashed_key - child.balance_limit = 75_000 - integration_session.add(parent) - integration_session.add(child) - await integration_session.commit() - - sent_requests: list[httpx.Request] = [] - - async def fake_transport( - request: httpx.Request, *args: Any, **kwargs: Any - ) -> httpx.Response: - sent_requests.append(request) - return _upstream_response(request) - - with ( - patch( - "httpx.AsyncHTTPTransport.handle_async_request", - side_effect=fake_transport, - ), - patch( - "routstr.payment.cost_calculation.sats_usd_price", - return_value=0.0005, - ), - ): - response = await authenticated_client.post( - "/v1/chat/completions", - json={ - "model": "dual-model", - "messages": [{"role": "user", "content": "hello"}], - }, - ) - - assert response.status_code == 200 - # The 100-sat candidate is rejected before forwarding; the third serves. - assert [request.url.host for request in sent_requests] == [ - "cheap.example.com", - "third.example.com", - ] - - await integration_session.refresh(parent) - await integration_session.refresh(child) - assert parent.reserved_balance == 0 - assert child.reserved_balance == 0 - assert parent.total_spent == response.json()["cost"]["total_msats"] - - records = ( - await integration_session.exec( - select(ReservationRelease).where(ReservationRelease.key_hash == key_hash) - ) - ).all() - assert len(records) == 2 - assert sorted(record.status for record in records) == ["charged", "released"] - assert len({record.reserved_msats for record in records}) == 1 - assert all(record.status != "active" for record in records) - - @pytest.fixture async def raised_envelope_provider_maps( patched_db_engine: None, diff --git a/tests/integration/test_free_response_stale_reservation.py b/tests/integration/test_free_response_stale_reservation.py index dadb6a6a..4b0b3c93 100644 --- a/tests/integration/test_free_response_stale_reservation.py +++ b/tests/integration/test_free_response_stale_reservation.py @@ -7,7 +7,7 @@ from unittest.mock import patch import pytest from sqlmodel.ext.asyncio.session import AsyncSession -from routstr.core.db import ApiKey +from routstr.core.db import ApiKey, ReservationRelease from routstr.payment.cost_calculation import CostData @@ -34,20 +34,30 @@ def _cost_data(total_msats: int) -> CostData: @pytest.mark.asyncio -async def test_overrun_charges_after_reservation_swept( +async def test_overrun_with_corrupted_aggregate_releases_without_charging( integration_session: AsyncSession, ) -> None: - """Overrun finalize must charge even when the reservation was already released.""" - from routstr.auth import adjust_payment_for_tokens, pay_for_request + """An overrun whose aggregate reservation was externally zeroed must not + charge: subtracting the reservation would have to clamp, which can erase + sibling reservations. The reservation is released and the charge dropped. + """ + from routstr import auth + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) deducted_max_cost = 990 # discounted reservation actual_token_cost = 1000 # actual cost overruns the reservation - # Sweeper has zeroed reserved_balance but left balance untouched. + # Something zeroed reserved_balance under an active durable reservation. key = _make_key(balance=1000, reserved=0) + key_hash = key.hashed_key integration_session.add(key) await integration_session.commit() await pay_for_request(key, deducted_max_cost, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) key.reserved_balance = 0 integration_session.add(key) await integration_session.commit() @@ -61,20 +71,67 @@ async def test_overrun_charges_after_reservation_swept( "routstr.auth.calculate_cost", return_value=_cost_data(actual_token_cost), ): - await adjust_payment_for_tokens( + result = await adjust_payment_for_tokens( key, response_data, integration_session, deducted_max_cost, None, None ) - await integration_session.refresh(key) + assert result["charged_msats"] == 0 + integration_session.expunge_all() + key_row = await integration_session.get(ApiKey, key_hash) + assert key_row is not None - assert key.total_spent == actual_token_cost, ( - f"Request was not billed (total_spent={key.total_spent}) — free response bug" + assert key_row.total_spent == 0, "corrupted aggregate must not be charged into" + assert key_row.balance == 1000 + assert key_row.reserved_balance == 0 + + # The corrupt reservation must reach a terminal state — an active leftover + # would be renewed by its heartbeat forever and poison stale cleanup. + record = await integration_session.get(ReservationRelease, reservation.release_id) + assert record is not None and record.status == "released" + assert reservation.release_id not in auth._reservation_heartbeats + + +@pytest.mark.asyncio +async def test_missing_usage_never_turns_reservation_into_charge( + integration_session: AsyncSession, +) -> None: + """A reservation is an authorization ceiling, not evidence of usage. + + Upstream handlers should provide locally estimated usage when possible. If + no measurement or estimate reaches settlement, release the reservation + rather than charging its full value. + """ + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, ) - assert key.balance == 1000 - actual_token_cost, ( - f"Balance not charged: {key.balance}" + + reserved = 4_000 + key = _make_key(balance=10_000, reserved=0) + key_hash = key.hashed_key + integration_session.add(key) + await integration_session.commit() + await pay_for_request(key, reserved, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + # No `usage` key at all — the upstream stream dropped its final usage chunk. + response_data = {"model": "test-model"} + result = await adjust_payment_for_tokens( + key, + response_data, + integration_session, + reserved, + reservation_snapshot=reservation, ) - assert key.balance >= 0 - assert key.reserved_balance == 0 + + assert result["charged_msats"] == 0 + integration_session.expunge_all() + key_row = await integration_session.get(ApiKey, key_hash) + assert key_row is not None + assert key_row.total_spent == 0 + assert key_row.balance == 10_000 + assert key_row.reserved_balance == 0 @pytest.mark.asyncio diff --git a/tests/integration/test_key_logic.py b/tests/integration/test_key_logic.py index 79d2c9fe..0a0fd573 100644 --- a/tests/integration/test_key_logic.py +++ b/tests/integration/test_key_logic.py @@ -1,14 +1,10 @@ -import asyncio import time -from datetime import datetime, timedelta import pytest -from fastapi import HTTPException -from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession from routstr.auth import pay_for_request -from routstr.core.db import ApiKey, create_session +from routstr.core.db import ApiKey @pytest.mark.asyncio @@ -25,367 +21,6 @@ async def test_key_validity_date(integration_session: AsyncSession) -> None: assert "expired" in str(excinfo.value).lower() -@pytest.mark.asyncio -async def test_key_balance_limit(integration_session: AsyncSession) -> None: - # 1. Create a key with a balance limit - key = ApiKey( - hashed_key="limited_key", balance=10000, balance_limit=500, total_spent=450 - ) - integration_session.add(key) - await integration_session.commit() - - # 2. Try to pay for a request that exceeds the limit - with pytest.raises(Exception) as excinfo: - await pay_for_request(key, 100, integration_session) - assert "limit exceeded" in str(excinfo.value).lower() - - # 3. Try to pay for a request that fits - await pay_for_request(key, 50, integration_session) - await integration_session.refresh(key) - # Note: total_spent is updated in adjust_payment_for_tokens, - # but pay_for_request checks it. - # In our current logic, pay_for_request checks (total_spent + cost) > balance_limit. - - -@pytest.mark.asyncio -async def test_key_daily_reset_policy(integration_session: AsyncSession) -> None: - # 1. Create a key with a daily reset policy and old reset date - yesterday = int((datetime.now() - timedelta(days=1)).timestamp()) - key = ApiKey( - hashed_key="daily_reset_key", - balance=10000, - balance_limit=1000, - balance_limit_reset="daily", - balance_limit_reset_date=yesterday, - total_spent=900, - ) - integration_session.add(key) - await integration_session.commit() - - # 2. Pay for a request - should trigger reset first because it's a new day - # Request is 200, total_spent is 900. 900+200 > 1000, - # but reset should happen making total_spent 0, then 0+200 < 1000. - await pay_for_request(key, 200, integration_session) - - await integration_session.refresh(key) - assert key.total_spent == 0 # Reset in pay_for_request happens before charging - # Wait, the charging logic in pay_for_request increments parent/billing_key's total_requests, - # but total_spent is updated in adjust_payment_for_tokens. - # However, the reset logic sets total_spent to 0. - assert key.balance_limit_reset_date is not None - assert key.balance_limit_reset_date > yesterday - - -@pytest.mark.asyncio -async def test_periodic_key_reset_job(integration_session: AsyncSession) -> None: - # 1. Create multiple keys needing reset - yesterday = int((datetime.now() - timedelta(days=1)).timestamp()) - key1 = ApiKey( - hashed_key="job_reset_key_1", - balance=1000, - balance_limit=1000, - balance_limit_reset="daily", - balance_limit_reset_date=yesterday, - total_spent=500, - ) - key2 = ApiKey( - hashed_key="job_reset_key_2", - balance=1000, - balance_limit=1000, - balance_limit_reset="daily", - balance_limit_reset_date=yesterday, - total_spent=800, - ) - integration_session.add(key1) - integration_session.add(key2) - await integration_session.commit() - - # 2. Run the periodic reset logic manually (mocking the background task loop) - # We can't easily run the actual loop because it has a sleep, - # but we can test the logic inside. - - # Implementation of periodic_key_reset logic for testing: - stmt = select(ApiKey).where(ApiKey.balance_limit_reset != None) # noqa: E711 - keys = (await integration_session.exec(stmt)).all() - now = int(time.time()) - for k in keys: - if k.hashed_key in ["job_reset_key_1", "job_reset_key_2"]: - k.total_spent = 0 - k.balance_limit_reset_date = now - integration_session.add(k) - await integration_session.commit() - - # 3. Verify resets - await integration_session.refresh(key1) - await integration_session.refresh(key2) - assert key1.total_spent == 0 - assert key2.total_spent == 0 - - -@pytest.mark.asyncio -async def test_balance_limit_enforced_atomically_under_concurrency( - patched_db_engine: None, -) -> None: - parent_hash = "parent_limit_atomic" - child_hash = "child_limit_atomic" - cost = 300 - - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=10000) - child = ApiKey( - hashed_key=child_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=cost, - total_spent=0, - ) - session.add(parent) - session.add(child) - await session.commit() - - results: list[str] = [] - - async def attempt() -> None: - async with create_session() as session: - fresh_child = await session.get(ApiKey, child_hash) - assert fresh_child is not None - try: - await pay_for_request(fresh_child, cost, session) - results.append("success") - except HTTPException as exc: - assert exc.status_code == 402 - results.append("blocked") - - await asyncio.gather(attempt(), attempt()) - - assert sorted(results) == ["blocked", "success"], ( - f"Expected exactly one success and one 402, got: {results}" - ) - - async with create_session() as session: - final_child = await session.get(ApiKey, child_hash) - assert final_child is not None - - assert final_child.reserved_balance == cost, ( - f"Child reserved_balance should equal one reservation, " - f"got {final_child.reserved_balance}" - ) - - -@pytest.mark.asyncio -async def test_parallel_payments_with_parent_and_child_key( - patched_db_engine: None, -) -> None: - parent_hash = "parent_parallel_mixed" - child_hash = "child_parallel_mixed" - cost = 300 - - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=10000) - child = ApiKey( - hashed_key=child_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=2 * cost, - ) - session.add(parent) - session.add(child) - await session.commit() - - async def attempt(key_hash: str) -> str: - async with create_session() as session: - fresh_key = await session.get(ApiKey, key_hash) - assert fresh_key is not None - try: - await pay_for_request(fresh_key, cost, session) - return "success" - except HTTPException as exc: - assert exc.status_code == 402 - return "blocked" - - results = await asyncio.gather(attempt(parent_hash), attempt(child_hash)) - - assert results == ["success", "success"], ( - f"Both parent and child payments should succeed, got: {results}" - ) - - async with create_session() as session: - final_parent = await session.get(ApiKey, parent_hash) - final_child = await session.get(ApiKey, child_hash) - assert final_parent is not None - assert final_child is not None - - # Both requests bill the parent; only the child request reserves on the child. - assert final_parent.reserved_balance == 2 * cost - assert final_parent.total_requests == 2 - assert final_child.reserved_balance == cost - assert final_child.total_requests == 1 - - -@pytest.mark.asyncio -async def test_balance_limit_with_existing_total_spent_under_concurrency( - patched_db_engine: None, -) -> None: - parent_hash = "parent_total_spent" - child_hash = "child_total_spent" - cost = 300 - - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=10000) - # 700 already spent against a 1000 limit: only one more 300 request fits. - child = ApiKey( - hashed_key=child_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=1000, - total_spent=700, - ) - session.add(parent) - session.add(child) - await session.commit() - - results: list[str] = [] - - async def attempt() -> None: - async with create_session() as session: - fresh_child = await session.get(ApiKey, child_hash) - assert fresh_child is not None - try: - await pay_for_request(fresh_child, cost, session) - results.append("success") - except HTTPException as exc: - assert exc.status_code == 402 - results.append("blocked") - - await asyncio.gather(attempt(), attempt()) - - assert sorted(results) == ["blocked", "success"], ( - f"Expected exactly one success and one 402, got: {results}" - ) - - async with create_session() as session: - final_child = await session.get(ApiKey, child_hash) - assert final_child is not None - - assert final_child.reserved_balance == cost - assert final_child.total_spent == 700 - - -@pytest.mark.asyncio -async def test_balance_limit_with_existing_reserved_balance( - patched_db_engine: None, -) -> None: - parent_hash = "parent_reserved_set" - blocked_hash = "child_reserved_blocked" - allowed_hash = "child_reserved_allowed" - cost = 300 - - async with create_session() as session: - parent = ApiKey(hashed_key=parent_hash, balance=10000) - # 800 already reserved against a 1000 limit: another 300 must be rejected. - blocked_child = ApiKey( - hashed_key=blocked_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=1000, - reserved_balance=800, - ) - # 500 reserved against a 1000 limit: another 300 still fits. - allowed_child = ApiKey( - hashed_key=allowed_hash, - balance=0, - parent_key_hash=parent_hash, - balance_limit=1000, - reserved_balance=500, - ) - session.add(parent) - session.add(blocked_child) - session.add(allowed_child) - await session.commit() - - async with create_session() as session: - fresh_blocked = await session.get(ApiKey, blocked_hash) - assert fresh_blocked is not None - with pytest.raises(HTTPException) as exc_info: - await pay_for_request(fresh_blocked, cost, session) - assert exc_info.value.status_code == 402 - - async with create_session() as session: - fresh_allowed = await session.get(ApiKey, allowed_hash) - assert fresh_allowed is not None - await pay_for_request(fresh_allowed, cost, session) - - async with create_session() as session: - final_blocked = await session.get(ApiKey, blocked_hash) - final_allowed = await session.get(ApiKey, allowed_hash) - final_parent = await session.get(ApiKey, parent_hash) - assert final_blocked is not None - assert final_allowed is not None - assert final_parent is not None - - assert final_blocked.reserved_balance == 800, "Rejected request must not reserve" - assert final_blocked.total_requests == 0 - assert final_allowed.reserved_balance == 500 + cost - assert final_allowed.total_requests == 1 - # Only the allowed request should have billed the parent. - assert final_parent.reserved_balance == cost - assert final_parent.total_requests == 1 - - -@pytest.mark.asyncio -async def test_child_reservation_discarded_when_parent_balance_depleted( - patched_db_engine: None, -) -> None: - parent_hash = "parent_depleted" - child_hash = "child_depleted" - cost = 300 - - async with create_session() as session: - # Parent can only afford one request; child has no balance_limit. - parent = ApiKey(hashed_key=parent_hash, balance=cost) - child = ApiKey( - hashed_key=child_hash, - balance=0, - parent_key_hash=parent_hash, - ) - session.add(parent) - session.add(child) - await session.commit() - - results: list[str] = [] - - async def attempt() -> None: - async with create_session() as session: - fresh_child = await session.get(ApiKey, child_hash) - assert fresh_child is not None - try: - await pay_for_request(fresh_child, cost, session) - results.append("success") - except HTTPException as exc: - assert exc.status_code == 402 - results.append("blocked") - - await asyncio.gather(attempt(), attempt()) - - assert sorted(results) == ["blocked", "success"], ( - f"Expected exactly one success and one 402, got: {results}" - ) - - async with create_session() as session: - final_parent = await session.get(ApiKey, parent_hash) - final_child = await session.get(ApiKey, child_hash) - assert final_parent is not None - assert final_child is not None - - assert final_parent.reserved_balance == cost - # The failed request must not leave a committed reservation on the child. - assert final_child.reserved_balance == cost, ( - f"Child reserved_balance should reflect only the successful request, " - f"got {final_child.reserved_balance}" - ) - assert final_child.total_requests == 1 - - @pytest.mark.asyncio async def test_refund_does_not_delete_key(integration_session: AsyncSession) -> None: # This requires mocking the router call or testing the logic in balance.py diff --git a/tests/integration/test_lightning_invoice_constraints.py b/tests/integration/test_lightning_invoice_constraints.py index 1a6d94b9..7370ebe4 100644 --- a/tests/integration/test_lightning_invoice_constraints.py +++ b/tests/integration/test_lightning_invoice_constraints.py @@ -1,10 +1,10 @@ """Integration tests for Lightning invoice key constraint fields. Covers two things: -- The three constraint fields (balance_limit, balance_limit_reset, validity_date) - are persisted on LightningInvoice and survive a DB round-trip. -- The production-path API-key record helper propagates those fields to the - created ApiKey, so the constraints are actually enforced when the key is used. +- The validity_date constraint field is persisted on LightningInvoice and + survives a DB round-trip. +- The production-path API-key record helper propagates it to the created + ApiKey, so the constraint is actually enforced when the key is used. """ from __future__ import annotations @@ -67,32 +67,6 @@ def mock_wallet_mint() -> object: # --------------------------------------------------------------------------- -@pytest.mark.asyncio -async def test_invoice_persists_balance_limit( - integration_session: AsyncSession, -) -> None: - invoice = _make_invoice(balance_limit=5000) - integration_session.add(invoice) - await integration_session.commit() - - stored = await integration_session.get(LightningInvoice, invoice.id) - assert stored is not None - assert stored.balance_limit == 5000 - - -@pytest.mark.asyncio -async def test_invoice_persists_balance_limit_reset( - integration_session: AsyncSession, -) -> None: - invoice = _make_invoice(balance_limit=5000, balance_limit_reset="daily") - integration_session.add(invoice) - await integration_session.commit() - - stored = await integration_session.get(LightningInvoice, invoice.id) - assert stored is not None - assert stored.balance_limit_reset == "daily" - - @pytest.mark.asyncio async def test_invoice_persists_validity_date( integration_session: AsyncSession, @@ -112,38 +86,6 @@ async def test_invoice_persists_validity_date( # --------------------------------------------------------------------------- -@pytest.mark.asyncio -async def test_created_key_receives_balance_limit( - integration_session: AsyncSession, -) -> None: - invoice = _make_invoice(balance_limit=8000) - integration_session.add(invoice) - await integration_session.flush() - - api_key = await _create_api_key_record(invoice, integration_session) - await integration_session.commit() - - stored_key = await integration_session.get(ApiKey, api_key.hashed_key) - assert stored_key is not None - assert stored_key.balance_limit == 8000 - - -@pytest.mark.asyncio -async def test_created_key_receives_balance_limit_reset( - integration_session: AsyncSession, -) -> None: - invoice = _make_invoice(balance_limit=8000, balance_limit_reset="monthly") - integration_session.add(invoice) - await integration_session.flush() - - api_key = await _create_api_key_record(invoice, integration_session) - await integration_session.commit() - - stored_key = await integration_session.get(ApiKey, api_key.hashed_key) - assert stored_key is not None - assert stored_key.balance_limit_reset == "monthly" - - @pytest.mark.asyncio async def test_created_key_receives_validity_date( integration_session: AsyncSession, @@ -422,8 +364,6 @@ async def test_created_key_without_constraints_has_none_fields( stored_key = await integration_session.get(ApiKey, api_key.hashed_key) assert stored_key is not None - assert stored_key.balance_limit is None - assert stored_key.balance_limit_reset is None assert stored_key.validity_date is None diff --git a/tests/integration/test_model_price_propagation.py b/tests/integration/test_model_price_propagation.py new file mode 100644 index 00000000..d637bbc8 --- /dev/null +++ b/tests/integration/test_model_price_propagation.py @@ -0,0 +1,221 @@ +"""Cover that an admin price write reaches the served catalogue (GET /v1/models). + +A price edit changes the served price; the served price is fee-adjusted while the +admin read-back is raw; a disabled model leaves the catalogue but keeps its row. + +Each test reaches the served model by id rather than iterating ``data["data"]``, +so an empty catalogue fails these tests instead of skipping them. +""" + +from __future__ import annotations + +from collections.abc import Iterator +from datetime import datetime, timedelta, timezone +from unittest.mock import patch + +import pytest +from httpx import AsyncClient +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import admin_sessions +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.proxy import reinitialize_upstreams + + +# The conftest patches ``routstr.payment.price.sats_usd_price``, but +# ``models.py`` imports it as ``from .price import sats_usd_price`` — a +# local binding the conftest-level patch cannot reach. Pin it here so +# every test that goes through ``_row_to_model`` gets a real sats price. +@pytest.fixture(autouse=True) +def _pin_sats_usd() -> Iterator[None]: + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + yield + + +def _admin_headers() -> dict[str, str]: + token = "test-propagation-token" + admin_sessions[token] = int( + (datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp() + ) + return {"Authorization": f"Bearer {token}"} + + +def _model_payload(prompt: float, provider_id: int, enabled: bool = True) -> dict: + return { + "id": "propagation-test-model", + "name": "Propagation Test Model", + "description": "model used to verify price propagation", + "created": 0, + "context_length": 128000, + "architecture": { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + }, + "pricing": { + "prompt": prompt, + "completion": prompt * 2, + "input_cache_read": 0.0, + "input_cache_write": 0.0, + "request": 0.0, + "image": 0.0, + "web_search": 0.0, + "internal_reasoning": 0.0, + }, + "per_request_limits": None, + "top_provider": None, + "upstream_provider_id": provider_id, + "canonical_slug": None, + "alias_ids": [], + "enabled": enabled, + "forwarded_model_id": "propagation-test-model", + } + + +async def _seed_provider(session: AsyncSession, *, fee: float = 1.0) -> int: + """Insert a provider, refresh the upstream map, and return its primary key.""" + provider = UpstreamProviderRow( + provider_type="generic", + base_url="https://propagation-test.example/v1", + api_key="test-key", + provider_fee=fee, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + await reinitialize_upstreams() + assert provider.id is not None + return provider.id + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_price_edit_propagates_to_served_catalogue( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """A price edit through the admin API must change the served /v1/models price.""" + + provider_id = await _seed_provider(integration_session) + + headers = _admin_headers() + r = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=headers, + json=_model_payload(prompt=1.0e-7, provider_id=provider_id), + ) + assert r.status_code == 200 + + # -- record the served price before edit ----------------------------------- + public = await integration_client.get("/v1/models") + assert public.status_code == 200 + public_data = public.json() + assert len(public_data["data"]) > 0, "catalogue must not be empty" + served_before = { + m["id"]: m.get("pricing", {}).get("prompt") for m in public_data["data"] + } + assert "propagation-test-model" in served_before + before = served_before["propagation-test-model"] + + # -- edit the price and re-check ------------------------------------------- + r = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=headers, + json=_model_payload(prompt=5.0e-7, provider_id=provider_id), + ) + assert r.status_code == 200 + + public = await integration_client.get("/v1/models") + assert public.status_code == 200 + served_after = { + m["id"]: m.get("pricing", {}).get("prompt") for m in public.json()["data"] + } + after = served_after["propagation-test-model"] + + assert before != after, "served price did not change after admin edit" + # With provider_fee=1.0 the served price equals the stored raw price. + assert after == pytest.approx(5.0e-7) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_readback_is_raw_served_is_fee_adjusted( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """Admin read-back returns the raw price; /v1/models returns the fee-adjusted one.""" + + provider_id = await _seed_provider(integration_session, fee=1.05) + + headers = _admin_headers() + model_payload = _model_payload(prompt=1.0e-7, provider_id=provider_id) + r = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=headers, + json=model_payload, + ) + assert r.status_code == 200 + + # Admin read-back: apply_provider_fee=False + admin_r = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/models/propagation-test-model", + headers=headers, + ) + assert admin_r.status_code == 200 + admin_body = admin_r.json() + raw_prompt = admin_body["pricing"]["prompt"] + assert raw_prompt == pytest.approx(1.0e-7) + + # Public /v1/models: fee-adjusted + public = await integration_client.get("/v1/models") + assert public.status_code == 200 + served = { + m["id"]: m.get("pricing", {}).get("prompt") for m in public.json()["data"] + } + assert served["propagation-test-model"] == pytest.approx(1.0e-7 * 1.05) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_disabled_model_not_served( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """A disabled model must be absent from /v1/models but still present in the DB.""" + + provider_id = await _seed_provider(integration_session) + + headers = _admin_headers() + r = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=headers, + json=_model_payload(prompt=1.0e-7, provider_id=provider_id, enabled=True), + ) + assert r.status_code == 200 + + # Confirm it appears in the public catalogue. + public = await integration_client.get("/v1/models") + served_ids = {m["id"] for m in public.json()["data"]} + assert "propagation-test-model" in served_ids + + # -- disable via upsert ---------------------------------------------------- + r = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=headers, + json=_model_payload(prompt=1.0e-7, provider_id=provider_id, enabled=False), + ) + assert r.status_code == 200 + + # Public catalogue must no longer list it. + public = await integration_client.get("/v1/models") + served_ids = {m["id"] for m in public.json()["data"]} + assert "propagation-test-model" not in served_ids + + # DB row must still exist. + row = await integration_session.get( + ModelRow, ("propagation-test-model", provider_id) + ) + assert row is not None + assert row.enabled is False diff --git a/tests/integration/test_model_serialization.py b/tests/integration/test_model_serialization.py new file mode 100644 index 00000000..791669d7 --- /dev/null +++ b/tests/integration/test_model_serialization.py @@ -0,0 +1,435 @@ +"""Characterization tests for the model-serialisation pipeline ``_row_to_model`` runs. + +These pin what the three public surfaces that reach it serve today: ``GET +/v1/models``, the admin single-model read-back, and the admin provider model +listing. Covered are the plain model, litellm cache backfill, preservation of +explicit cache rates, the request-price floor, the provider-fee flag against +recomputed max costs, survival of a failed sats conversion, the full serialised +dict field for field, and agreement between the two admin views. + +Nothing here asserts on the deterministic USD half in isolation, so the pins hold +whether or not it is split out from the live BTC-rate conversion. +""" + +from __future__ import annotations + +import json +from collections.abc import Iterator +from datetime import datetime, timedelta, timezone +from unittest.mock import patch + +import pytest +from httpx import AsyncClient +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import admin_sessions +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.proxy import reinitialize_upstreams + + +# The conftest patches ``routstr.payment.price.sats_usd_price``, but +# ``models.py`` imports it as ``from .price import sats_usd_price`` — a +# local binding the conftest-level patch cannot reach. Pin it here so +# every test that goes through ``_row_to_model`` gets a real sats price. +@pytest.fixture(autouse=True) +def _pin_sats_usd() -> Iterator[None]: + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + yield + + +def _admin_headers() -> dict[str, str]: + token = "test-serialisation-token" + admin_sessions[token] = int( + (datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp() + ) + return {"Authorization": f"Bearer {token}"} + + +# -- helpers ------------------------------------------------------------------- + +_SEEDED_MODEL_ID = "ser-test-model" + + +async def _seed_provider(session: AsyncSession, fee: float = 1.0) -> int: + """Insert a provider, refresh the upstream map, and return its primary key.""" + provider = UpstreamProviderRow( + provider_type="generic", + base_url="https://serialisation-test.example/v1", + api_key="test-key", + provider_fee=fee, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + await reinitialize_upstreams() + assert provider.id is not None + return provider.id + + +async def _seed_model( + session: AsyncSession, + provider_id: int, + *, + model_id: str = _SEEDED_MODEL_ID, + prompt: float = 1.0e-7, + completion: float = 2.0e-7, + cache_read: float = 0.0, + cache_write: float = 0.0, + request_price: float = 0.0, + enabled: bool = True, +) -> ModelRow: + row = ModelRow( + id=model_id, + name=f"SerTest {model_id}", + description="characterization model", + created=0, + context_length=128000, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + } + ), + pricing=json.dumps( + { + "prompt": prompt, + "completion": completion, + "input_cache_read": cache_read, + "input_cache_write": cache_write, + "request": request_price, + "image": 0.0, + "web_search": 0.0, + "internal_reasoning": 0.0, + } + ), + upstream_provider_id=provider_id, + enabled=enabled, + forwarded_model_id=model_id, + ) + session.add(row) + await session.commit() + return row + + +async def _raw_via_admin(client: AsyncClient, provider_id: int, model_id: str) -> dict: + """Return the raw (``apply_provider_fee=False``) model dict from admin read-back.""" + r = await client.get( + f"/admin/api/upstream-providers/{provider_id}/models/{model_id}", + headers=_admin_headers(), + ) + assert r.status_code == 200 + return r.json() + + +async def _served_via_public(client: AsyncClient, model_id: str) -> dict | None: + """Return the served model dict from /v1/models, or None if absent.""" + r = await client.get("/v1/models") + assert r.status_code == 200 + return {m["id"]: m for m in r.json()["data"]}.get(model_id) + + +# -- test 1: plain model ------------------------------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_plain_model_serialisation( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """A model with prompt + completion prices has the expected serialised shape.""" + provider_id = await _seed_provider(integration_session) + await _seed_model(integration_session, provider_id) + await reinitialize_upstreams() + + # Admin read-back (raw, no fee) + body = await _raw_via_admin( + integration_client, + provider_id, + _SEEDED_MODEL_ID, + ) + assert body["id"] == _SEEDED_MODEL_ID + assert body["pricing"]["prompt"] == pytest.approx(1.0e-7) + assert body["pricing"]["completion"] == pytest.approx(2.0e-7) + assert body["sats_pricing"] is not None + # With sats_usd_price = 0.0005 (the fixture) + assert body["sats_pricing"]["prompt"] == pytest.approx(1.0e-7 / 0.0005) + assert body["sats_pricing"]["completion"] == pytest.approx(2.0e-7 / 0.0005) + + # Public /v1/models (fee applied) + s = await _served_via_public(integration_client, _SEEDED_MODEL_ID) + assert s is not None, f"{_SEEDED_MODEL_ID} not found in /v1/models" + # fee=1.0 so values match raw + assert s["pricing"]["prompt"] == pytest.approx(1.0e-7) + assert s["pricing"]["completion"] == pytest.approx(2.0e-7) + + +# -- test 2: no cache rates (litellm backfill) --------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_without_cache_rates_gets_litellm_backfill( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """A well-known model without cache rates gets them from litellm's cost map.""" + provider_id = await _seed_provider(integration_session) + # Use a real litellm-known id so backfill_cache_pricing can find it. + await _seed_model( + integration_session, + provider_id, + model_id="gpt-4o", + prompt=2.5e-6, + completion=1.0e-5, + cache_read=0.0, + cache_write=0.0, + ) + await reinitialize_upstreams() + + # Admin read-back (raw): cache_read should be present after backfill. + # (cache_write may not be in litellm's map for every model.) + body = await _raw_via_admin( + integration_client, + provider_id, + "gpt-4o", + ) + assert body["pricing"]["input_cache_read"] > 0.0, ( + "backfill_cache_pricing should have filled input_cache_read from litellm" + ) + + +# -- test 3: cache rates already present (not overwritten) --------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_existing_cache_rates_not_overwritten( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """A model with explicit cache rates must keep them; the backfill is a no-op.""" + provider_id = await _seed_provider(integration_session) + await _seed_model( + integration_session, + provider_id, + prompt=1.0e-7, + completion=2.0e-7, + cache_read=9.99e-9, + cache_write=8.88e-9, + ) + await reinitialize_upstreams() + + body = await _raw_via_admin( + integration_client, + provider_id, + _SEEDED_MODEL_ID, + ) + assert body["pricing"]["input_cache_read"] == pytest.approx(9.99e-9) + assert body["pricing"]["input_cache_write"] == pytest.approx(8.88e-9) + + +# -- test 4: request price floor ----------------------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_with_request_price( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """A model with a request price floor carries it through to the served model.""" + provider_id = await _seed_provider(integration_session) + await _seed_model(integration_session, provider_id, request_price=0.01) + await reinitialize_upstreams() + + body = await _raw_via_admin( + integration_client, + provider_id, + _SEEDED_MODEL_ID, + ) + assert body["pricing"]["request"] == pytest.approx(0.01) + + +# -- test 5: provider_fee=True vs False, max costs recomputed ------------------ + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_fee_flag_changes_pricing_but_max_costs_are_recomputed( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """With provider_fee=1.5, fee-adjusted pricing is 1.5× raw, but max costs are NOT.""" + provider_id = await _seed_provider(integration_session, fee=1.5) + await _seed_model(integration_session, provider_id) + await reinitialize_upstreams() + + # Admin read-back: raw, no fee. + body = await _raw_via_admin( + integration_client, + provider_id, + _SEEDED_MODEL_ID, + ) + assert body["pricing"]["prompt"] == pytest.approx(1.0e-7) + + # Public /v1/models: fee applied. + s = await _served_via_public(integration_client, _SEEDED_MODEL_ID) + assert s is not None + assert s["pricing"]["prompt"] == pytest.approx(1.0e-7 * 1.5) + + # max_prompt_cost must NOT just be multiplied by 1.5 — it is recomputed from + # the fee-inflated per-token rates and context_length. + cl = 128_000 + expected_max_prompt = cl * 1.0e-7 * 1.5 + assert s["pricing"]["max_prompt_cost"] == pytest.approx(expected_max_prompt) + + +# -- test 6: sats conversion failure keeps the model alive --------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_model_survives_sats_conversion_failure( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """When the BTC feed fails, the model still returns — with no sats_pricing.""" + provider_id = await _seed_provider(integration_session) + await _seed_model(integration_session, provider_id) + + # Make sats_usd_price raise so _update_model_sats_pricing swallows it. + # The admin read-back must happen inside the patch block. + with patch( + "routstr.payment.models.sats_usd_price", + side_effect=RuntimeError("BTC feed down"), + ): + await reinitialize_upstreams() + body = await _raw_via_admin( + integration_client, + provider_id, + _SEEDED_MODEL_ID, + ) + assert body["id"] == _SEEDED_MODEL_ID + assert body["sats_pricing"] is None, ( + "sats conversion failure must not crash — the model returns with no sats_pricing" + ) + + +# -- test 7: the whole serialised dict ----------------------------------------- + +# The sats figures are the USD ones divided by the pinned 0.0005 rate; written +# as the division so the expectation carries the same float error the code does. +_SATS = 0.0005 + + +def _expected_serialised_model(provider_id: int) -> dict: + """Every field the raw admin read-back produces for the seeded model. + + Note the max-cost asymmetry: with fees off the USD max costs stay at zero + while their sats counterparts are computed. That is what the code does + today, and pinning it is the point. + """ + return { + "alias_ids": None, + "architecture": { + "input_modalities": ["text"], + "instruct_type": None, + "modality": "text", + "output_modalities": ["text"], + "tokenizer": "unknown", + }, + "canonical_slug": None, + "context_length": 128000, + "created": 0, + "description": "characterization model", + "enabled": True, + "forwarded_model_id": _SEEDED_MODEL_ID, + "id": _SEEDED_MODEL_ID, + "name": f"SerTest {_SEEDED_MODEL_ID}", + "per_request_limits": None, + "pricing": { + "completion": 2.0e-7, + "image": 0.0, + "input_cache_read": 0.0, + "input_cache_write": 0.0, + "internal_reasoning": 0.0, + "max_completion_cost": 0.0, + "max_cost": 0.0, + "max_prompt_cost": 0.0, + "prompt": 1.0e-7, + "request": 0.01, + "web_search": 0.0, + }, + "sats_pricing": { + "completion": 2.0e-7 / _SATS, + "image": 0.0, + "input_cache_read": 0.0, + "input_cache_write": 0.0, + "internal_reasoning": 0.0, + "max_completion_cost": 0.0, + "max_cost": 0.001, + "max_prompt_cost": 0.0, + "prompt": 1.0e-7 / _SATS, + "request": 0.01 / _SATS, + "web_search": 0.0, + }, + "top_provider": None, + "upstream_provider_id": provider_id, + } + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_full_serialised_model_is_unchanged( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """Pin every field of the serialised model, not just the interesting ones. + + The tests above pin values. A refactor that dropped a field outright + would satisfy all of them and fail only here. + """ + provider_id = await _seed_provider(integration_session) + await _seed_model(integration_session, provider_id, request_price=0.01) + await reinitialize_upstreams() + + body = await _raw_via_admin( + integration_client, + provider_id, + _SEEDED_MODEL_ID, + ) + assert body == _expected_serialised_model(provider_id) + + +# -- test 8: the provider model listing ---------------------------------------- + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_provider_listing_matches_single_model_read_back( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """The listing is a third entry into the same builder — it must agree.""" + provider_id = await _seed_provider(integration_session) + await _seed_model(integration_session, provider_id, request_price=0.01) + await reinitialize_upstreams() + + single = await _raw_via_admin( + integration_client, + provider_id, + _SEEDED_MODEL_ID, + ) + + r = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + ) + assert r.status_code == 200 + listed = {m["id"]: m for m in r.json()["db_models"]} + assert _SEEDED_MODEL_ID in listed, "seeded model missing from the provider listing" + assert listed[_SEEDED_MODEL_ID] == single diff --git a/tests/integration/test_negative_available_balance_repro.py b/tests/integration/test_negative_available_balance_repro.py new file mode 100644 index 00000000..b31769f1 --- /dev/null +++ b/tests/integration/test_negative_available_balance_repro.py @@ -0,0 +1,747 @@ +"""Regression tests for the production "negative available balance" 402. + +Invariants protected: +* billing and the admin API agree on what "available" means +* an in-flight (heartbeaten) reservation is never swept; an abandoned one is + released and can never be charged afterwards +* a cost overrun spends only its own reservation plus unreserved balance +* corrupt reservations are repaired terminally instead of poisoning cleanup +* under concurrency, balance / reserved / available never go negative +""" + +import asyncio +import random +import time +import uuid +from typing import Awaitable, Callable +from unittest.mock import patch + +import pytest +from httpx import AsyncClient +from sqlmodel import col, func, select, update +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey, ReservationRelease +from routstr.payment.cost_calculation import CostData + +pytestmark = pytest.mark.integration + +# Realistic sweeper timeout: a renewed (heartbeaten) reservation stays alive, +# a reservation backdated past this is released. +STALE_TIMEOUT_SECONDS = 300 + +# created_at is whole seconds; -1 makes every reservation stale immediately. +# Only used in the fuzz test to stress the terminal-release path. +SWEEP_EVERYTHING = -1 + + +def _cost_data(total_msats: int) -> CostData: + return CostData( + base_msats=0, + input_msats=total_msats // 2, + output_msats=total_msats - total_msats // 2, + total_msats=total_msats, + total_usd=0.0, + input_tokens=100, + output_tokens=100, + ) + + +def _response(model: str = "test-model") -> dict: + return { + "model": model, + "usage": {"prompt_tokens": 100, "completion_tokens": 100}, + } + + +async def _new_key(session: AsyncSession, balance: int) -> str: + key_hash = f"test_neg_{uuid.uuid4().hex}" + session.add( + ApiKey( + hashed_key=key_hash, + balance=balance, + reserved_balance=0, + total_spent=0, + total_requests=0, + ) + ) + await session.commit() + return key_hash + + +async def _backdate_reservation( + session: AsyncSession, release_id: str, seconds: int +) -> None: + """Age a reservation's lease as if it had not been renewed for `seconds`.""" + await session.exec( # type: ignore[call-overload] + update(ReservationRelease) + .where(col(ReservationRelease.id) == release_id) + .values(created_at=col(ReservationRelease.created_at) - seconds) + ) + await session.commit() + + +async def _wait_for( + predicate: "Callable[[], Awaitable[bool]]", + timeout: float = 10.0, + interval: float = 0.1, +) -> bool: + """Bounded polling instead of fixed sleeps for background-task effects.""" + deadline = asyncio.get_event_loop().time() + timeout + while True: + if await predicate(): + return True + if asyncio.get_event_loop().time() > deadline: + return False + await asyncio.sleep(interval) + + +@pytest.mark.asyncio +async def test_402_reports_negative_available_while_admin_shows_positive( + integration_session: AsyncSession, + integration_client: AsyncClient, +) -> None: + """The production symptom: billing rejects on balance - reserved_balance, + so the admin endpoint must expose reserved/available, not just balance.""" + from fastapi import HTTPException + + from routstr.auth import _validate_bearer_key_locked + from routstr.core.db import set_admin_password + + key_hash = await _new_key(integration_session, balance=263_000) + # Leaked reservations slightly exceeding the balance, as in production. + await integration_session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == key_hash) + .values(reserved_balance=267_215) + ) + await integration_session.commit() + + with pytest.raises(HTTPException) as exc: + await _validate_bearer_key_locked( + "sk-" + key_hash, integration_session, min_cost=1 + ) + + assert exc.value.status_code == 402 + message = exc.value.detail["error"]["message"] # type: ignore[index] + assert "-4.215 sats (-4215 msats) available" in message, message + + await set_admin_password(integration_session, "test-admin-pw") + login = await integration_client.post( + "/admin/api/login", json={"password": "test-admin-pw"} + ) + assert login.status_code == 200 + token = login.json()["token"] + resp = await integration_client.get( + "/admin/api/temporary-balances", + params={"search": key_hash}, + headers={"Authorization": f"Bearer {token}"}, + ) + assert resp.status_code == 200 + rows = [row for row in resp.json()["balances"] if row["hashed_key"] == key_hash] + assert len(rows) == 1 + row = rows[0] + assert row["balance"] == 263_000 + assert row["reserved_balance"] == 267_215 + assert row["available_balance"] == -4_215 + assert resp.json()["totals"] == { + "total_balance": 263_000, + "total_reserved_balance": 267_215, + "total_available_balance": -4_215, + "total_spent": 0, + "total_requests": 0, + } + + +@pytest.mark.asyncio +async def test_abandoned_reservation_is_swept_and_cannot_finalize( + integration_session: AsyncSession, +) -> None: + """Release is terminal: a swept reservation's late finalizer must not + charge, or it could spend funds since reserved by another request.""" + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + from routstr.core.db import release_stale_reservations + + cost = 5_000 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, cost, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + # Client vanished: the lease is never renewed and ages past the timeout. + await _backdate_reservation( + integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 1 + ) + released = await release_stale_reservations( + integration_session, STALE_TIMEOUT_SECONDS + ) + assert released == 1 + + # A zombie finalizer shows up afterwards; it must not charge. + with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)): + await adjust_payment_for_tokens( + key, + _response(), + integration_session, + cost, + reservation_snapshot=reservation, + ) + + await integration_session.refresh(key) + record = await integration_session.get(ReservationRelease, reservation.release_id) + assert record is not None and record.status == "released" + assert key.total_spent == 0 + assert key.balance == 10_000 + assert key.reserved_balance == 0 + + +@pytest.mark.asyncio +async def test_sweeper_cannot_release_a_reservation_renewed_after_selection( + integration_session: AsyncSession, +) -> None: + """A heartbeat landing between the sweeper's select and its transition + must win — exercised through the public sweeper entry point.""" + from routstr.auth import ( + get_reservation_snapshot, + pay_for_request, + renew_reservation, + ) + from routstr.core import db as core_db + from routstr.core.db import release_stale_reservations + + key_hash = await _new_key(integration_session, balance=1_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + await pay_for_request(key, 1_000, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + await _backdate_reservation( + integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 100 + ) + + real_transition = core_db._transition_stale_reservation + + async def renew_then_transition( + session: AsyncSession, reservation_id: str, cutoff: int + ) -> bool: + # The sweeper selected this reservation as stale; the heartbeat + # renews exactly between that select and the transition. + assert await renew_reservation(reservation, session) + return await real_transition(session, reservation_id, cutoff) + + with patch.object(core_db, "_transition_stale_reservation", renew_then_transition): + released = await release_stale_reservations( + integration_session, STALE_TIMEOUT_SECONDS + ) + + assert released == 0 + record = await integration_session.get(ReservationRelease, reservation.release_id) + assert record is not None and record.status == "active" + await integration_session.refresh(key) + assert key.reserved_balance == 1_000 + + +@pytest.mark.asyncio +async def test_legacy_cleanup_cannot_erase_a_reservation_committed_after_its_read( + integration_session: AsyncSession, +) -> None: + """A reservation committing between the legacy sweep's read and its + zeroing must survive — exercised through the public sweeper entry point.""" + from routstr.core import db as core_db + from routstr.core.db import release_stale_reservations + + key_hash = await _new_key(integration_session, balance=10_000) + stale_reserved_at = int(time.time()) - (STALE_TIMEOUT_SECONDS + 100) + await integration_session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == key_hash) + .values(reserved_balance=500, reserved_at=stale_reserved_at) + ) + await integration_session.commit() + + real_release = core_db._release_legacy_aggregate + + async def commit_reservation_then_release( + session: AsyncSession, + target_key_hash: str, + observed_reserved: int, + observed_reserved_at: int | None, + ) -> bool: + # The sweeper read the legacy aggregate and found no active durable + # owner; a new reservation commits exactly before the zeroing lands. + session.add( + ReservationRelease( + id=uuid.uuid4().hex, + key_hash=target_key_hash, + billing_key_hash=target_key_hash, + reserved_msats=700, + status="active", + ) + ) + await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == target_key_hash) + .values( + reserved_balance=col(ApiKey.reserved_balance) + 700, + reserved_at=int(time.time()), + ) + ) + await session.commit() + return await real_release( + session, target_key_hash, observed_reserved, observed_reserved_at + ) + + with patch.object( + core_db, "_release_legacy_aggregate", commit_reservation_then_release + ): + released = await release_stale_reservations( + integration_session, STALE_TIMEOUT_SECONDS + ) + + assert released == 0 + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + assert key.reserved_balance == 1_200, "legacy cleanup erased a live reservation" + + +@pytest.mark.asyncio +async def test_sweeper_repairs_corrupt_reservation_and_continues_batch( + integration_session: AsyncSession, +) -> None: + """One corrupt durable reservation (aggregate no longer holds its msats) + must be terminalized without aggregate subtraction, and must not stop the + rest of the batch from being released normally.""" + from routstr.auth import get_reservation_snapshot, pay_for_request + from routstr.core.db import release_stale_reservations + + cost = 1_000 + corrupt_hash = await _new_key(integration_session, balance=cost) + healthy_hash = await _new_key(integration_session, balance=cost) + + corrupt_key = await integration_session.get(ApiKey, corrupt_hash) + assert corrupt_key is not None + await pay_for_request(corrupt_key, cost, integration_session) + corrupt_reservation = await get_reservation_snapshot( + corrupt_key, integration_session + ) + + healthy_key = await integration_session.get(ApiKey, healthy_hash) + assert healthy_key is not None + await pay_for_request(healthy_key, cost, integration_session) + healthy_reservation = await get_reservation_snapshot( + healthy_key, integration_session + ) + + # Corrupt the first key: its aggregate no longer holds the reservation. + await integration_session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == corrupt_hash) + .values(reserved_balance=0) + ) + await integration_session.commit() + + for reservation in (corrupt_reservation, healthy_reservation): + await _backdate_reservation( + integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 100 + ) + + released = await release_stale_reservations( + integration_session, STALE_TIMEOUT_SECONDS + ) + assert released == 2, "a corrupt record must not abort the sweep batch" + + for reservation in (corrupt_reservation, healthy_reservation): + record = await integration_session.get( + ReservationRelease, reservation.release_id + ) + assert record is not None and record.status == "released" + healthy_key = await integration_session.get(ApiKey, healthy_hash) + corrupt_key = await integration_session.get(ApiKey, corrupt_hash) + assert healthy_key is not None and corrupt_key is not None + assert healthy_key.reserved_balance == 0 + assert corrupt_key.reserved_balance == 0 + assert corrupt_key.balance == cost, "repair must not touch balances" + + +@pytest.mark.asyncio +async def test_heartbeat_survives_a_rolled_back_charge_attempt( + integration_session: AsyncSession, + patched_db_engine: None, +) -> None: + """Claiming a reservation must not stop its heartbeat: a rollback restores + the active reservation, which then still needs lease renewal.""" + from routstr import auth + from routstr.auth import ( + _claim_reservation_for_charge, + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + from routstr.core.db import create_session + + cost = 1_000 + timeout = 3 # heartbeat interval = 1s + with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout): + async with create_session() as session: + key_hash = await _new_key(session, balance=2_000) + key = await session.get(ApiKey, key_hash) + assert key is not None + await pay_for_request(key, cost, session) + reservation = await get_reservation_snapshot(key, session) + + try: + # A charge attempt claims the reservation, then its transaction + # fails and rolls back. + async with create_session() as session: + assert await _claim_reservation_for_charge(reservation, session) + await session.rollback() + + async with create_session() as session: + record = await session.get(ReservationRelease, reservation.release_id) + assert record is not None and record.status == "active" + await _backdate_reservation( + session, reservation.release_id, timeout * 10 + ) + record = await session.get(ReservationRelease, reservation.release_id) + assert record is not None + backdated_lease = record.created_at + + async def lease_renewed() -> bool: + async with create_session() as session: + record = await session.get( + ReservationRelease, reservation.release_id + ) + return record is not None and record.created_at > backdated_lease + + assert await _wait_for(lease_renewed), ( + "heartbeat did not survive the rolled-back charge attempt" + ) + + # The restored reservation still finalizes normally. + with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)): + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + result = await adjust_payment_for_tokens( + key, + _response(), + session, + cost, + reservation_snapshot=reservation, + ) + finally: + await auth._stop_reservation_heartbeat(reservation.release_id) + + assert result["charged_msats"] == cost + async with create_session() as session: + record = await session.get(ReservationRelease, reservation.release_id) + assert record is not None and record.status == "charged" + key = await session.get(ApiKey, key_hash) + assert key is not None + assert key.total_spent == cost + + +@pytest.mark.asyncio +async def test_heartbeat_dies_with_its_request_so_sweeper_can_recover( + integration_session: AsyncSession, + patched_db_engine: None, +) -> None: + """A request that vanishes without finalizing must not renew forever — + its heartbeat stops with the owning task and the sweeper reclaims the + funds.""" + from routstr import auth + from routstr.auth import get_reservation_snapshot, pay_for_request + from routstr.core.db import create_session, release_stale_reservations + + cost = 1_000 + timeout = 3 # heartbeat interval = 1s + async with create_session() as session: + key_hash = await _new_key(session, balance=cost) + + holder: dict = {} + with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout): + + async def doomed_request() -> None: + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + await pay_for_request(key, cost, session) + holder["reservation"] = await get_reservation_snapshot(key, session) + # ...request control dies here, no finalize and no release. + + await asyncio.create_task(doomed_request()) + release_id = holder["reservation"].release_id + + try: + # While the heartbeat is still winding down it may renew once + # more; keep backdating until the sweeper wins, which it must as + # soon as the dead owner is noticed. + async def sweeper_recovered() -> bool: + async with create_session() as session: + await _backdate_reservation(session, release_id, timeout * 10) + return await release_stale_reservations(session, timeout) == 1 + + assert await _wait_for(sweeper_recovered), ( + "sweeper never recovered the abandoned reservation" + ) + finally: + await auth._stop_reservation_heartbeat(release_id) + + async with create_session() as session: + record = await session.get(ReservationRelease, release_id) + assert record is not None and record.status == "released" + key = await session.get(ApiKey, key_hash) + assert key is not None + assert key.reserved_balance == 0 + assert key.total_spent == 0 + + +@pytest.mark.asyncio +async def test_reservation_heartbeat_covers_the_whole_request_lifecycle( + integration_session: AsyncSession, + patched_db_engine: None, +) -> None: + """pay_for_request starts the heartbeat, finalization stops it, and a + backdated lease is renewed in the background without any manual call.""" + from routstr import auth + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + from routstr.core.db import create_session, release_stale_reservations + + cost = 1_000 + timeout = 3 # heartbeat interval = 1s + async with create_session() as session: + key_hash = await _new_key(session, balance=2 * cost) + + with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout): + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + await pay_for_request(key, cost, session) + reservation = await get_reservation_snapshot(key, session) + + try: + # Only a background renewal can keep this alive now. + async with create_session() as session: + await _backdate_reservation( + session, reservation.release_id, timeout * 10 + ) + record = await session.get(ReservationRelease, reservation.release_id) + assert record is not None + backdated_lease = record.created_at + + async def lease_renewed() -> bool: + async with create_session() as session: + record = await session.get( + ReservationRelease, reservation.release_id + ) + return record is not None and record.created_at > backdated_lease + + assert await _wait_for(lease_renewed), "heartbeat never renewed the lease" + async with create_session() as session: + assert await release_stale_reservations(session, timeout) == 0 + + with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)): + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + result = await adjust_payment_for_tokens( + key, + _response(), + session, + cost, + reservation_snapshot=reservation, + ) + finally: + await auth._stop_reservation_heartbeat(reservation.release_id) + + assert result["charged_msats"] == cost + async with create_session() as session: + record = await session.get(ReservationRelease, reservation.release_id) + assert record is not None and record.status == "charged" + key = await session.get(ApiKey, key_hash) + assert key is not None + assert key.total_spent == cost + assert key.reserved_balance == 0 + + +@pytest.mark.asyncio +async def test_overrun_cannot_spend_a_concurrent_reservation( + integration_session: AsyncSession, +) -> None: + """An overrun is capped to its own reservation plus unreserved balance; + the sibling's reserved funds stay untouched and available stays >= 0.""" + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + + reserved_each = 100 + overrun_cost = 150 # A's real token cost exceeds its reservation + + # Balance covers exactly two reservations; nothing free on top. + key_hash = await _new_key(integration_session, balance=2 * reserved_each) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, reserved_each, integration_session) + reservation_a = await get_reservation_snapshot(key, integration_session) + await pay_for_request(key, reserved_each, integration_session) + await get_reservation_snapshot(key, integration_session) # B stays in flight + + await integration_session.refresh(key) + assert key.reserved_balance == 2 * reserved_each + + # Only A finalizes; B is still streaming and its funds must stay reserved. + with patch("routstr.auth.calculate_cost", return_value=_cost_data(overrun_cost)): + result = await adjust_payment_for_tokens( + key, + _response(), + integration_session, + reserved_each, + reservation_snapshot=reservation_a, + ) + + assert result["charged_msats"] == reserved_each + assert result["total_msats"] == overrun_cost + await integration_session.refresh(key) + assert key.total_spent == reserved_each + assert key.reserved_balance == reserved_each + assert key.balance == reserved_each + assert key.total_balance >= 0, ( + f"available balance went negative: balance={key.balance} " + f"reserved={key.reserved_balance} -> {key.total_balance} msats; " + "the overrun charge consumed the still-reserved funds of request B" + ) + + +@pytest.mark.asyncio +async def test_concurrent_requests_with_sweeper_keep_balance_invariants( + integration_session: AsyncSession, + patched_db_engine: None, +) -> None: + """Fuzz: concurrent requests against an everything-is-stale sweeper. + Some requests legitimately finish uncharged (release is terminal), but + balances must never go negative and no reservation may stay active.""" + from fastapi import HTTPException + + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + from routstr.core.db import create_session, release_stale_reservations + + rng = random.Random(1337) + starting_balance = 200_000 + n_requests = 24 + + async with create_session() as session: + key_hash = await _new_key(session, balance=starting_balance) + + completed_costs: list[int] = [] + rejected_requests = 0 + + async def one_request(index: int) -> None: + nonlocal rejected_requests + reserved = rng.randrange(1_000, 4_000) + actual = max(1, int(reserved * rng.uniform(0.5, 1.1))) + try: + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + await pay_for_request(key, reserved, session) + except HTTPException as exc: + # A depleted balance is the only legitimate rejection. + assert exc.status_code == 402, exc.detail + rejected_requests += 1 + return + + try: + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + reservation = await get_reservation_snapshot(key, session) + except RuntimeError: + # The everything-is-stale sweeper can release the reservation + # before the stream even starts; the request aborts uncharged. + return + + await asyncio.sleep(rng.uniform(0, 0.02)) # the "stream" + + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + await adjust_payment_for_tokens( + key, + _response(str(actual)), + session, + reserved, + reservation_snapshot=reservation, + ) + completed_costs.append(actual) + + sweeping = True + + async def sweeper() -> None: + while sweeping: + async with create_session() as session: + await release_stale_reservations(session, SWEEP_EVERYTHING) + await asyncio.sleep(0.002) + + sweep_task = asyncio.create_task(sweeper()) + try: + with patch( + "routstr.auth.calculate_cost", + side_effect=lambda response_data, *a, **k: _cost_data( + int(response_data["model"]) + ), + ): + await asyncio.gather(*(one_request(i) for i in range(n_requests))) + finally: + sweeping = False + await sweep_task + + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + leftover_active = ( + await session.exec( # type: ignore[call-overload] + select(func.count()) + .select_from(ReservationRelease) + .where(col(ReservationRelease.status) == "active") + ) + ).one() + + assert completed_costs or rejected_requests, "no request made any progress" + assert key.balance >= 0, f"balance went negative: {key.balance}" + assert key.reserved_balance >= 0, ( + f"reserved_balance went negative: {key.reserved_balance}" + ) + assert key.total_balance >= 0, ( + f"available balance negative: balance={key.balance} " + f"reserved={key.reserved_balance} (this is the production symptom)" + ) + assert key.total_spent <= starting_balance, ( + f"spent {key.total_spent} of a {starting_balance} balance" + ) + # Requests whose reservation was swept mid-flight finish uncharged, so the + # charged total can only be at most the sum of completed request costs. + max_expected_spend = sum(completed_costs) + assert key.total_spent <= max_expected_spend, ( + f"charged {key.total_spent} msats but completed requests only cost " + f"{max_expected_spend}" + ) + assert leftover_active == 0, ( + f"{leftover_active} reservations still active after all requests finished" + ) diff --git a/tests/integration/test_payment_invariants.py b/tests/integration/test_payment_invariants.py new file mode 100644 index 00000000..3d5460dd --- /dev/null +++ b/tests/integration/test_payment_invariants.py @@ -0,0 +1,450 @@ +"""Invariant coverage for the reserve → charge → release money path. + +Every finalization branch of ``adjust_payment_for_tokens`` must respect the +same accounting rules: a completed request is charged exactly once, its +reported ``charged_msats`` matches the actual debit, it never spends more than +its own reservation leaves available, and concurrent requests cannot raid each +other's reservations. +""" + +import asyncio +import uuid +from unittest.mock import patch + +import pytest +from sqlmodel import col, select, update +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey, ReservationRelease +from routstr.payment.cost_calculation import ( + CostData, + CostDataError, + MaxCostData, +) + +pytestmark = pytest.mark.integration + + +def _cost_data(total_msats: int, cls: type[CostData] = CostData) -> CostData: + return cls( + base_msats=0, + input_msats=total_msats // 2, + output_msats=total_msats - total_msats // 2, + total_msats=total_msats, + total_usd=0.0, + input_tokens=100, + output_tokens=100, + ) + + +def _response() -> dict: + return { + "model": "test-model", + "usage": {"prompt_tokens": 100, "completion_tokens": 100}, + } + + +async def _new_key(session: AsyncSession, balance: int) -> str: + key_hash = f"test_inv_{uuid.uuid4().hex}" + session.add( + ApiKey( + hashed_key=key_hash, + balance=balance, + reserved_balance=0, + total_spent=0, + total_requests=0, + ) + ) + await session.commit() + return key_hash + + +async def _active_reservations(session: AsyncSession) -> int: + rows = await session.exec( + select(ReservationRelease).where(col(ReservationRelease.status) == "active") + ) + return len(rows.all()) + + +@pytest.mark.asyncio +async def test_corrupt_revert_still_decrements_request_count( + integration_session: AsyncSession, +) -> None: + from routstr.auth import ( + get_reservation_snapshot, + pay_for_request, + revert_pay_for_request, + ) + + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + await pay_for_request(key, 3_000, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + await integration_session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == key_hash) + .values(reserved_balance=0) + ) + await integration_session.commit() + + assert await revert_pay_for_request( + key, + integration_session, + 3_000, + reservation_snapshot=reservation, + ) + updated = await integration_session.get(ApiKey, key_hash) + record = await integration_session.get(ReservationRelease, reservation.release_id) + assert updated is not None + assert updated.total_requests == 0 + assert updated.reserved_balance == 0 + assert record is not None and record.status == "released" + + +@pytest.mark.asyncio +async def test_exact_cost_branch_charges_the_reservation_once( + integration_session: AsyncSession, +) -> None: + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + + cost = 3_000 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, cost, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)): + result = await adjust_payment_for_tokens( + key, + _response(), + integration_session, + cost, + reservation_snapshot=reservation, + ) + + assert result["charged_msats"] == cost + await integration_session.refresh(key) + assert key.balance == 10_000 - cost + assert key.total_spent == cost + assert key.reserved_balance == 0 + assert key.total_requests == 1 + assert await _active_reservations(integration_session) == 0 + + +@pytest.mark.asyncio +async def test_max_cost_branch_charges_the_reservation_once( + integration_session: AsyncSession, +) -> None: + """No token pricing configured -> flat MaxCostData charge of the reservation.""" + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + + cost = 2_500 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, cost, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + with patch( + "routstr.auth.calculate_cost", + return_value=_cost_data(cost, cls=MaxCostData), + ): + result = await adjust_payment_for_tokens( + key, + _response(), + integration_session, + cost, + reservation_snapshot=reservation, + ) + + assert result["charged_msats"] == cost + await integration_session.refresh(key) + assert key.balance == 10_000 - cost + assert key.total_spent == cost + assert key.reserved_balance == 0 + + +@pytest.mark.asyncio +async def test_underrun_branch_refunds_the_unused_reservation( + integration_session: AsyncSession, +) -> None: + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + + reserved = 5_000 + actual = 1_200 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, reserved, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + with patch("routstr.auth.calculate_cost", return_value=_cost_data(actual)): + result = await adjust_payment_for_tokens( + key, + _response(), + integration_session, + reserved, + reservation_snapshot=reservation, + ) + + assert result["charged_msats"] == actual + await integration_session.refresh(key) + assert key.total_spent == actual, "user must pay the real cost, not the reservation" + assert key.balance == 10_000 - actual + assert key.reserved_balance == 0 + + +@pytest.mark.asyncio +async def test_overrun_branch_charges_full_cost_when_balance_is_free( + integration_session: AsyncSession, +) -> None: + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + + reserved = 1_000 + actual = 1_400 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, reserved, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + with patch("routstr.auth.calculate_cost", return_value=_cost_data(actual)): + result = await adjust_payment_for_tokens( + key, + _response(), + integration_session, + reserved, + reservation_snapshot=reservation, + ) + + assert result["charged_msats"] == actual + await integration_session.refresh(key) + assert key.total_spent == actual + assert key.balance == 10_000 - actual + assert key.reserved_balance == 0 + + +@pytest.mark.asyncio +async def test_zero_cost_response_is_free_and_releases_the_reservation( + integration_session: AsyncSession, +) -> None: + """An empty/unusable upstream response must cost the user nothing.""" + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + + reserved = 4_000 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, reserved, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + with patch("routstr.auth.calculate_cost", return_value=_cost_data(0)): + await adjust_payment_for_tokens( + key, + {"model": "test-model"}, + integration_session, + reserved, + reservation_snapshot=reservation, + ) + + await integration_session.refresh(key) + assert key.balance == 10_000 + assert key.total_spent == 0 + assert key.reserved_balance == 0 + assert await _active_reservations(integration_session) == 0 + + +@pytest.mark.asyncio +async def test_cost_error_releases_the_reservation_without_charging( + integration_session: AsyncSession, +) -> None: + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + + reserved = 4_000 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, reserved, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + with patch( + "routstr.auth.calculate_cost", + return_value=CostDataError(message="no pricing", code="pricing_error"), + ): + cost = await adjust_payment_for_tokens( + key, + _response(), + integration_session, + reserved, + reservation_snapshot=reservation, + ) + assert cost["charged_msats"] == 0 + + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + assert key.balance == 10_000, "a pricing failure must not charge the user" + assert key.total_spent == 0 + assert key.reserved_balance == 0, "funds must not stay locked" + assert await _active_reservations(integration_session) == 0 + + +@pytest.mark.asyncio +async def test_repeated_finalization_charges_only_once( + integration_session: AsyncSession, +) -> None: + """A retried finalizer (proxy retry, duplicate stream end) must not double-bill.""" + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + + cost = 3_000 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, cost, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + results = [] + with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)): + for _ in range(3): + # A declined re-charge rolls its session back, so re-load the key + # the way a fresh request would instead of reusing a stale instance. + integration_session.expunge_all() + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + results.append( + await adjust_payment_for_tokens( + key, + _response(), + integration_session, + cost, + reservation_snapshot=reservation, + ) + ) + + # Only the first finalization debits; duplicates report a zero charge. + assert [r["charged_msats"] for r in results] == [cost, 0, 0] + + integration_session.expunge_all() + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + assert key.total_spent == cost, f"charged {key.total_spent} for one request" + assert key.balance == 10_000 - cost + + +@pytest.mark.asyncio +async def test_concurrent_duplicate_finalization_charges_only_once( + integration_session: AsyncSession, + patched_db_engine: None, +) -> None: + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + ) + from routstr.core.db import create_session + + cost = 3_000 + async with create_session() as session: + key_hash = await _new_key(session, balance=10_000) + key = await session.get(ApiKey, key_hash) + assert key is not None + await pay_for_request(key, cost, session) + reservation = await get_reservation_snapshot(key, session) + + async def finalize() -> None: + async with create_session() as session: + fresh = await session.get(ApiKey, key_hash) + assert fresh is not None + await adjust_payment_for_tokens( + fresh, + _response(), + session, + cost, + reservation_snapshot=reservation, + ) + + with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)): + await asyncio.gather(finalize(), finalize(), finalize()) + + async with create_session() as session: + key = await session.get(ApiKey, key_hash) + assert key is not None + assert key.total_spent == cost + assert key.balance == 10_000 - cost + assert key.reserved_balance == 0 + + +@pytest.mark.asyncio +async def test_release_after_charge_does_not_credit_the_user_back( + integration_session: AsyncSession, +) -> None: + """A late cleanup path must not turn a charged request into a free one.""" + from routstr.auth import ( + adjust_payment_for_tokens, + get_reservation_snapshot, + pay_for_request, + release_reservation, + ) + + cost = 3_000 + key_hash = await _new_key(integration_session, balance=10_000) + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + + await pay_for_request(key, cost, integration_session) + reservation = await get_reservation_snapshot(key, integration_session) + + with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)): + await adjust_payment_for_tokens( + key, + _response(), + integration_session, + cost, + reservation_snapshot=reservation, + ) + + released = await release_reservation(reservation, integration_session, cost) + assert released is False, "a charged reservation must not be releasable" + + key = await integration_session.get(ApiKey, key_hash) + assert key is not None + assert key.balance == 10_000 - cost + assert key.total_spent == cost + assert key.reserved_balance == 0 diff --git a/tests/integration/test_ppq_auto_topup_claim.py b/tests/integration/test_ppq_auto_topup_claim.py index 433389cd..05b99c03 100644 --- a/tests/integration/test_ppq_auto_topup_claim.py +++ b/tests/integration/test_ppq_auto_topup_claim.py @@ -23,6 +23,7 @@ from routstr.upstream.auto_topup import ( _ppq_request_id, _ppq_spent_last_24h_usd, _ppq_state_id_for_provider, + _reconcile_ppq_state, _record_ppq_invoice, _set_ppq_state_terminal, get_ppq_auto_topup_state, @@ -35,6 +36,8 @@ pytestmark = pytest.mark.asyncio def _row(provider_id: int = 1) -> MagicMock: row = MagicMock() row.id = provider_id + row.api_key = "secret" + row.provider_settings = None return row @@ -111,12 +114,35 @@ async def test_claim_is_reusable_once_the_previous_attempt_finished( await _seed_provider() first = await _claim_ppq_topup(_row()) assert first is not None - assert await _set_ppq_state_terminal(_row(), first, collected=True, swept=False) + assert await _set_ppq_state_terminal(_row(), first, collected=False, swept=True) second = await _claim_ppq_topup(_row()) assert second is not None and second != first +async def test_settled_claim_suppresses_immediate_duplicate( + patched_db_engine: Any, +) -> None: + await _seed_provider() + row = _row() + operation_id = await _claim_ppq_topup(row) + assert operation_id is not None + assert await _set_ppq_state_terminal(row, operation_id, collected=True, swept=False) + + assert await _reconcile_ppq_state(row, provider=None) is True + assert await _claim_ppq_topup(row) is None + + +async def test_claim_rejects_stale_provider_configuration( + patched_db_engine: Any, +) -> None: + await _seed_provider() + stale = _row() + stale.provider_settings = '{"auto_topup":true}' + + assert await _claim_ppq_topup(stale) is None + + async def test_recording_the_invoice_moves_the_claim_in_flight( patched_db_engine: Any, ) -> None: @@ -287,7 +313,13 @@ async def test_ppq_payment_audit_row_is_visible_and_survives_next_claim( assert audit["collected"] is True assert "lnbc-secret-invoice" not in audit["token"] - # Reusing the deterministic claim lock must not overwrite history. + assert await _claim_ppq_topup(_row()) is None + async with create_session() as session: + state = await session.get(CashuTransaction, _ppq_state_id_for_provider(1)) + assert state is not None + state.created_at = int(time.time()) - 301 + session.add(state) + await session.commit() assert await _claim_ppq_topup(_row()) is not None async with create_session() as session: assert await session.get(CashuTransaction, audit["id"]) is not None diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index b7db0c6c..de9a3112 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -686,7 +686,7 @@ async def test_no_database_changes_during_provider_operations( @pytest.mark.integration @pytest.mark.asyncio -async def test_admin_routstr_topup_retries_transient_upstream_failure( +async def test_admin_routstr_topup_does_not_duplicate_invoice_on_upstream_failure( integration_client: AsyncClient, integration_session: Any, ) -> None: @@ -739,16 +739,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure( assert json["api_key"] == "sk-upstream-test" assert headers["Authorization"] == "Bearer sk-upstream-test" - if self.calls == 1: - return MockResponse(500, {"detail": "warmup failure"}) - - return MockResponse( - 200, - { - "bolt11": "lnbc1testinvoice", - "invoice_id": "invoice-123", - }, - ) + return MockResponse(500, {"detail": "ambiguous upstream failure"}) mock_client = MockAsyncClient() @@ -759,11 +750,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure( json={"amount": 10}, ) - assert response.status_code == 200 - data = response.json() - assert data["ok"] is True - assert data["topup_data"]["payment_request"] == "lnbc1testinvoice" - assert data["topup_data"]["invoice_id"] == "invoice-123" - assert mock_client.calls == 2 + assert response.status_code == 500 + assert mock_client.calls == 1 finally: admin_sessions.pop(admin_token, None) diff --git a/tests/integration/test_proxy_post_endpoints.py b/tests/integration/test_proxy_post_endpoints.py index 8d5fab36..99bdc071 100644 --- a/tests/integration/test_proxy_post_endpoints.py +++ b/tests/integration/test_proxy_post_endpoints.py @@ -289,8 +289,115 @@ async def test_proxy_post_unauthorized_access(integration_client: AsyncClient) - assert response.status_code in [400, 401] +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "bad_path", + [ + "internal/admin", # unknown endpoint, no API prefix + "v1/../admin", # traversal onto a sibling path + "v1//models", # duplicate separator + "%2e%2e/secret", # encoded dot segment + ], +) +async def test_authenticated_post_to_unknown_path_is_rejected( + authenticated_client: AsyncClient, bad_path: str +) -> None: + """An authenticated POST to an unknown/traversal path must be rejected at + the edge (404) and never forwarded — the provider credential must not reach + an endpoint the caller merely spelled into the URL. If the guard let it + through, forwarding would raise and this would not be a clean 404.""" + with patch( + "routstr.upstream.base.BaseUpstreamProvider.forward_request", + AsyncMock(side_effect=AssertionError("must not forward unknown path")), + ): + response = await authenticated_client.post( + f"/{bad_path}", + json={"model": "gpt-3.5-turbo", "messages": []}, + ) + assert response.status_code == 404 +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "bad_path", + [ + # Unambiguously spelled, under a prefix the proxy serves, but not an + # endpoint it offers. These are real upstream routes that manage keys, + # org membership, and billing on the same origin as inference. + "v1/organization/api_keys", + "v1/api_keys", + "v1/billing/usage", + "v1/admin/keys", + # An id segment is honoured one level deep, and only where an endpoint + # takes one at all. + "v1/chat/completions/abc", + "v1/models/gpt-4/secret", + ], +) +async def test_known_prefix_does_not_carry_an_unknown_endpoint( + authenticated_client: AsyncClient, bad_path: str +) -> None: + """A familiar prefix must not be a passport for the rest of the origin. + + Nothing here is traversal-shaped, so the spelling screen lets it by; only + the endpoint allowlist stops it. Forwarding raises if the guard misses, + so a clean 404 also proves the credential never left.""" + with patch( + "routstr.upstream.base.BaseUpstreamProvider.forward_request", + AsyncMock(side_effect=AssertionError("must not forward unknown endpoint")), + ): + response = await authenticated_client.post( + f"/{bad_path}", + json={"model": "gpt-3.5-turbo", "messages": []}, + ) + assert response.status_code == 404 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_get_unknown_path_is_rejected(integration_client: AsyncClient) -> None: + """A GET to an unknown (no API prefix) path is rejected at the edge.""" + response = await integration_client.get("/internal/admin") + assert response.status_code == 404 + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "bad_path", + [ + # Percent-encoded so the traversal survives the client to the server, + # which decodes it to "v1/../admin" before routing. + "v1/%2e%2e/admin", + # Well-spelled, so only the endpoint allowlist can stop it. + "v1/organization/api_keys", + "anything/encrypted", + ], +) +async def test_ehbp_request_is_gated_by_the_endpoint_allowlist( + integration_client: AsyncClient, bad_path: str +) -> None: + """EHBP hides the body from the proxy, not the destination. + + The encrypted contract covers the request body; it says nothing about which + endpoint the provider credential gets spent against, so an EHBP request is + screened and allowlisted exactly like any other.""" + with patch( + "routstr.proxy.forward_ehbp_request", + AsyncMock(side_effect=AssertionError("must not forward unknown endpoint")), + ): + response = await integration_client.post( + f"/{bad_path}", + content=b"encrypted", + headers={ + "ehbp-encapsulated-key": "x", + "x-routstr-model": "gpt-4", + }, + ) + assert response.status_code == 404 + @pytest.mark.integration @pytest.mark.asyncio diff --git a/tests/integration/test_prune_dead_api_keys.py b/tests/integration/test_prune_dead_api_keys.py index de2b9d44..d4b54570 100644 --- a/tests/integration/test_prune_dead_api_keys.py +++ b/tests/integration/test_prune_dead_api_keys.py @@ -45,7 +45,7 @@ def _dead_key(created_at: int | None) -> ApiKey: @pytest.mark.asyncio async def test_prunes_old_refunded_zero_key(patched_db_engine: None) -> None: - """A funded-then-refunded key (0/0/0, NULL parent, old) is pruned.""" + """A funded-then-refunded zero-balance key with an old timestamp is pruned.""" key = _dead_key(LONG_AGO) async with create_session() as session: session.add(key) @@ -98,34 +98,6 @@ async def test_used_key_never_pruned(patched_db_engine: None) -> None: assert await _exists(k.hashed_key) -@pytest.mark.asyncio -async def test_parent_and_child_keys_are_not_pruned( - patched_db_engine: None, -) -> None: - """Pruning must not orphan child keys or delete valid children.""" - parent = _dead_key(LONG_AGO) - child = ApiKey( - hashed_key=f"child_{uuid.uuid4().hex}", - balance=0, - reserved_balance=0, - total_spent=0, - total_requests=0, - created_at=LONG_AGO, - parent_key_hash=parent.hashed_key, - ) - async with create_session() as session: - session.add(parent) - session.add(child) - await session.commit() - - async with create_session() as session: - pruned = await prune_dead_api_keys(session, OLD) - - assert pruned == 0 - assert await _exists(parent.hashed_key) - assert await _exists(child.hashed_key) - - @pytest.mark.asyncio @pytest.mark.parametrize( ("status", "expires_at"), @@ -321,3 +293,33 @@ async def test_periodic_prune_disabled_returns_immediately( await auth.periodic_dead_key_prune() sleep_mock.assert_not_called() + + +@pytest.mark.asyncio +async def test_refund_claim_protects_key(patched_db_engine: None) -> None: + """A key with a refund row is the audit anchor for its payout and holds a + non-null FK, so the janitor must leave it alone.""" + from routstr.core.db import Refund + + key = _dead_key(LONG_AGO) + async with create_session() as session: + session.add(key) + await session.commit() + session.add( + Refund( + api_key_hashed_key=key.hashed_key, + method="lightning", + destination="user@ln.example.com", + amount_msats=1000, + unit="sat", + mint_url="https://mint.example.com", + status="paid", + ) + ) + await session.commit() + + async with create_session() as session: + pruned = await prune_dead_api_keys(session, OLD) + + assert pruned == 0 + assert await _exists(key.hashed_key) diff --git a/tests/integration/test_refund_claims.py b/tests/integration/test_refund_claims.py new file mode 100644 index 00000000..cda5735e --- /dev/null +++ b/tests/integration/test_refund_claims.py @@ -0,0 +1,891 @@ +"""Refund claim lifecycle against a real SQLite database. + +Covers the guarantees the ``refunds`` table exists to provide: one open claim +per key, a persisted melt quote before the melt is dispatched, and a +reconciler that never restores a balance whose payout may have settled. +""" + +import time +from typing import Any, Awaitable, Callable +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +from fastapi import HTTPException +from sqlmodel import select + +from routstr import refund +from routstr.balance import RefundRequest, refund_wallet_endpoint +from routstr.core.db import ( + ApiKey, + AsyncSession, + CashuTransaction, + Refund, + store_cashu_transaction_with_retry, + total_user_liability, +) +from routstr.payment.lnurl import LNURLError, MeltOutcomeAmbiguousError + +KEY_HASH = "refundclaimkey" +ADDRESS = "user@ln.example.com" +BALANCE_MSATS = 5_000_000 + + +async def _seed_key( + session: AsyncSession, *, balance: int = BALANCE_MSATS, address: str | None = None +) -> ApiKey: + key = ApiKey(hashed_key=KEY_HASH) + key.balance = balance + key.reserved_balance = 0 + key.refund_currency = "sat" + key.refund_address = address + key.total_spent = 0 + key.total_requests = 0 + session.add(key) + await session.commit() + await session.refresh(key) + return key + + +async def _load_key(session: AsyncSession) -> ApiKey: + key = await session.get(ApiKey, KEY_HASH) + assert key is not None + await session.refresh(key) + return key + + +async def _load_refund(session: AsyncSession, refund_id: str) -> Refund: + row = await session.get(Refund, refund_id) + assert row is not None + await session.refresh(row) + return row + + +async def _age_claim(session: AsyncSession, refund_id: str, seconds: int) -> None: + row = await _load_refund(session, refund_id) + row.claimed_at = int(time.time()) - seconds + row.created_at = int(time.time()) - seconds + row.updated_at = int(time.time()) - seconds + session.add(row) + await session.commit() + + +@pytest.fixture +def short_timeout() -> Any: + with patch.object(refund.settings, "refund_claim_timeout_seconds", 300): + yield + + +# --- exclusivity ----------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_second_claim_on_open_key_is_rejected( + integration_session: AsyncSession, +) -> None: + key = await _seed_key(integration_session) + first = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + key = await _load_key(integration_session) + assert key.balance == 0 + + with pytest.raises(HTTPException) as exc_info: + await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + detail = exc_info.value.detail + assert exc_info.value.status_code == 409 + assert isinstance(detail, dict) + assert detail["error"]["code"] == "refund_in_progress" + + rows = (await integration_session.exec(select(Refund))).all() + assert [row.id for row in rows] == [first.id] + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_claim_from_second_session_hits_the_index( + integration_engine: Any, integration_session: AsyncSession +) -> None: + key = await _seed_key(integration_session) + await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + + async with AsyncSession(integration_engine, expire_on_commit=False) as other: + other_key = await _load_key(other) + with pytest.raises(HTTPException) as exc_info: + await refund.open_claim( + other, other_key, method="lightning", destination=ADDRESS + ) + assert exc_info.value.status_code == 409 + + +@pytest.mark.asyncio +async def test_claim_rejects_stale_balance_snapshot( + integration_session: AsyncSession, +) -> None: + key = await _seed_key(integration_session) + integration_session.expunge(key) # a stale, detached snapshot + key.balance = BALANCE_MSATS + 1 + with pytest.raises(HTTPException) as exc_info: + await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + assert exc_info.value.status_code == 409 + assert (await integration_session.exec(select(Refund))).all() == [] + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_retry_after_failed_claim_pays_once( + integration_session: AsyncSession, +) -> None: + key = await _seed_key(integration_session) + first = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + assert await refund.release(integration_session, first) + key = await _load_key(integration_session) + assert key.balance == BALANCE_MSATS + + second = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + assert second.id != first.id + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_release_after_settle_is_a_noop( + integration_session: AsyncSession, +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + assert await refund.settle(integration_session, claim, quote_id="q1") + assert not await refund.release(integration_session, claim) + assert (await _load_key(integration_session)).balance == 0 + row = await _load_refund(integration_session, claim.id) + assert row.status == "paid" + assert row.claimed_at is None + + +# --- execute --------------------------------------------------------------- + + +def _lnurl_stub( + outcome: BaseException | None = None, + *, + quoted: bool = True, +) -> Callable[..., Awaitable[int]]: + async def send( + amount: int, + unit: str, + mint: str, + address: str, + *, + on_melt_quote: Callable[[str, str], Awaitable[None]] | None = None, + ) -> int: + if quoted and on_melt_quote is not None: + await on_melt_quote("quote-123", mint) + if outcome is not None: + raise outcome + return amount + + return send + + +@pytest.mark.asyncio +async def test_execute_persists_quote_before_melt_and_settles( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + seen: list[str | None] = [] + + async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int: + await on_melt_quote("quote-123", claim.mint_url) + row = await _load_refund(integration_session, claim.id) + seen.append(row.quote_id) + return 5000 + + with patch("routstr.refund.send_to_lnurl", send): + body = await refund.execute(integration_session, claim) + + assert seen == ["quote-123"], "quote must be on disk before the melt runs" + assert body["status"] == "paid" + assert body["recipient"] == ADDRESS + assert body["sats"] == "5000" + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.quote_id, row.claimed_at) == ("paid", "quote-123", None) + + +@pytest.mark.asyncio +async def test_execute_ambiguous_holds_claim_with_quote( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + with patch( + "routstr.refund.send_to_lnurl", _lnurl_stub(MeltOutcomeAmbiguousError("?")) + ): + with pytest.raises(HTTPException) as exc_info: + await refund.execute(integration_session, claim) + + assert exc_info.value.status_code == 502 + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.quote_id, row.claimed_at) == ( + "ambiguous", + "quote-123", + None, + ) + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_execute_clean_failure_restores_balance( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + with patch( + "routstr.refund.send_to_lnurl", + _lnurl_stub(LNURLError("limits"), quoted=False), + ): + with pytest.raises(HTTPException) as exc_info: + await refund.execute(integration_session, claim) + + assert exc_info.value.status_code == 500 + row = await _load_refund(integration_session, claim.id) + assert row.status == "failed" + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_execute_failure_after_quote_withholds_balance( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """The mint may have paid the quote, so a later local failure must not restore.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + with patch("routstr.refund.send_to_lnurl", _lnurl_stub(RuntimeError("local"))): + with pytest.raises(HTTPException) as exc_info: + await refund.execute(integration_session, claim) + + assert exc_info.value.status_code == 502 + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.quote_id) == ("ambiguous", "quote-123") + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_execute_records_mint_that_issued_the_quote( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """Mint fallback must leave reconciliation pointed at the issuing mint.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + fallback_mint = "https://fallback.mint.example" + + async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int: + await on_melt_quote("quote-fallback", fallback_mint) + raise MeltOutcomeAmbiguousError("unknown") + + with patch("routstr.refund.send_to_lnurl", send): + with pytest.raises(MeltOutcomeAmbiguousError): + await refund._pay_lightning(integration_session, claim) + + row = await _load_refund(integration_session, claim.id) + assert (row.mint_url, row.quote_id, row.status) == ( + fallback_mint, + "quote-fallback", + "ambiguous", + ) + + +@pytest.mark.asyncio +async def test_execute_aborts_melt_when_claim_was_released( + integration_engine: Any, integration_session: AsyncSession, patched_db_engine: None +) -> None: + """A reconciler that released the claim first must stop the melt.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + melted = False + + async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int: + nonlocal melted + async with AsyncSession(integration_engine, expire_on_commit=False) as other: + await refund.release(other, await _load_refund(other, claim.id)) + await on_melt_quote("quote-123", claim.mint_url) + melted = True + return 5000 + + with patch("routstr.refund.send_to_lnurl", send): + with pytest.raises(HTTPException): + await refund.execute(integration_session, claim) + + assert melted is False + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +# --- reconciler ------------------------------------------------------------ + + +async def _open_ambiguous(session: AsyncSession, quote_id: str | None) -> Refund: + key = await _seed_key(session) + claim = await refund.open_claim( + session, key, method="lightning", destination=ADDRESS + ) + await refund.hold(session, claim, quote_id) + return claim + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("mint_status", "expected_status", "expected_balance"), + [ + ("paid", "paid", 0), + ("unpaid", "failed", BALANCE_MSATS), + ("pending", "ambiguous", 0), + ("unknown", "ambiguous", 0), + ], +) +async def test_reconcile_ambiguous_claims( + integration_session: AsyncSession, + patched_db_engine: None, + short_timeout: None, + mint_status: str, + expected_status: str, + expected_balance: int, +) -> None: + claim = await _open_ambiguous(integration_session, "quote-123") + await _age_claim(integration_session, claim.id, 600) + with patch( + "routstr.refund.check_bolt11_payment_status", + AsyncMock(return_value=mint_status), + ) as check: + await refund.reconcile_once() + + check.assert_awaited_once_with(claim.mint_url, "sat", "quote-123") + row = await _load_refund(integration_session, claim.id) + assert row.status == expected_status + assert (await _load_key(integration_session)).balance == expected_balance + + +@pytest.mark.asyncio +async def test_reconcile_credits_balance_once_across_passes( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + claim = await _open_ambiguous(integration_session, "quote-123") + await _age_claim(integration_session, claim.id, 600) + with patch( + "routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid") + ): + await refund.reconcile_once() + await refund.reconcile_once() + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_reconcile_waits_before_trusting_recent_unpaid( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + """A just-dispatched melt can report unpaid before it turns pending, so an + unpaid quote is only final once the claim has been quiet for a timeout.""" + claim = await _open_ambiguous(integration_session, "quote-123") + with patch( + "routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid") + ): + await refund.reconcile_once() + row = await _load_refund(integration_session, claim.id) + assert row.status == "ambiguous" + assert (await _load_key(integration_session)).balance == 0 + + await _age_claim(integration_session, claim.id, 600) + with patch( + "routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid") + ): + await refund.reconcile_once() + row = await _load_refund(integration_session, claim.id) + assert row.status == "failed" + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_reconcile_keeps_claim_that_gained_quote_mid_pass( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + """The payout records its quote between the reconciler's read and its + release: the stale ``quote_id is None`` must not restore the balance.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + await _age_claim(integration_session, claim.id, 600) + + real_lease = refund._lease + + async def lease_then_quote(refund_id: str, now: int, cutoff: int) -> bool: + leased = await real_lease(refund_id, now, cutoff) + await refund.record_quote(claim, "late-quote", claim.mint_url) + return leased + + with patch("routstr.refund._lease", lease_then_quote): + await refund.reconcile_once() + + row = await _load_refund(integration_session, claim.id) + assert row.status == "pending" + assert row.quote_id == "late-quote" + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_reconcile_leaves_fresh_pending_claim_alone( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + with patch("routstr.refund.check_bolt11_payment_status", AsyncMock()) as check: + await refund.reconcile_once() + check.assert_not_awaited() + row = await _load_refund(integration_session, claim.id) + assert row.status == "pending" + assert row.claimed_at is not None + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_reconcile_releases_expired_claim_without_quote( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + """No quote on disk means the mint was never asked to pay.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + await _age_claim(integration_session, claim.id, 600) + with patch("routstr.refund.check_bolt11_payment_status", AsyncMock()) as check: + await refund.reconcile_once() + check.assert_not_awaited() + assert (await _load_refund(integration_session, claim.id)).status == "failed" + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_reconcile_queries_mint_for_crashed_claim_with_quote( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + """Crash after the quote was persisted: the mint decides, not the timeout.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + await refund.record_quote(claim, "quote-crash", claim.mint_url) + await _age_claim(integration_session, claim.id, 600) + with patch( + "routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="paid") + ) as check: + await refund.reconcile_once() + check.assert_awaited_once_with(claim.mint_url, "sat", "quote-crash") + assert (await _load_refund(integration_session, claim.id)).status == "paid" + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_reconcile_marks_expired_cashu_claim_stuck( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + await _age_claim(integration_session, claim.id, 600) + with patch("routstr.refund.logger") as log: + await refund.reconcile_once() + await refund.reconcile_once() + assert log.critical.call_count == 1 + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.claimed_at) == ("stuck", None) + assert (await _load_key(integration_session)).balance == 0 + # A stuck claim is closed, so the key is not permanently locked out. + key = await _load_key(integration_session) + key.balance = 1000 + integration_session.add(key) + await integration_session.commit() + await refund.open_claim(integration_session, key, method="cashu", destination=None) + + +@pytest.mark.asyncio +async def test_reconcile_survives_one_failing_row( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + claim = await _open_ambiguous(integration_session, "quote-123") + with patch( + "routstr.refund.check_bolt11_payment_status", + AsyncMock(side_effect=RuntimeError("mint down")), + ): + await refund.reconcile_once() + row = await _load_refund(integration_session, claim.id) + assert row.status == "ambiguous" + assert row.claimed_at is not None, "lease is kept until the next pass" + + +# --- endpoint -------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_endpoint_uses_requested_address_over_persisted( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + await _seed_key(integration_session, address="stored@ln.example.com") + send = AsyncMock(side_effect=_lnurl_stub()) + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()) as resolve, + patch("routstr.refund.send_to_lnurl", send), + ): + body = await refund_wallet_endpoint( + refund_request=RefundRequest(lightning_address=ADDRESS), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + resolve.assert_awaited_once_with(ADDRESS) + assert isinstance(body, dict) + assert body["recipient"] == ADDRESS + assert send.await_args is not None + assert send.await_args.args[3] == ADDRESS + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_endpoint_rejects_bad_address_without_debit( + integration_session: AsyncSession, +) -> None: + await _seed_key(integration_session) + with patch( + "routstr.refund.get_lnurl_data", AsyncMock(side_effect=LNURLError("nope")) + ): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + refund_request=RefundRequest(lightning_address="bad@example"), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert exc_info.value.status_code == 400 + assert (await integration_session.exec(select(Refund))).all() == [] + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_endpoint_replays_paid_lightning_refund_on_empty_balance( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + await _seed_key(integration_session, address=ADDRESS) + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()), + patch("routstr.refund.send_to_lnurl", _lnurl_stub()), + ): + first = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + second = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(first, dict) and isinstance(second, dict) + assert second["refund_id"] == first["refund_id"] + assert second["status"] == "paid" + + +@pytest.mark.asyncio +async def test_endpoint_refund_while_ambiguous_returns_409( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + await _open_ambiguous(integration_session, "quote-123") + key = await _load_key(integration_session) + key.balance = 2_000_000 # topped up while the melt is unresolved + integration_session.add(key) + await integration_session.commit() + + send = AsyncMock() + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()), + patch("routstr.refund.send_to_lnurl", send), + ): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + refund_request=RefundRequest(lightning_address=ADDRESS), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert exc_info.value.status_code == 409 + send.assert_not_awaited() + assert (await _load_key(integration_session)).balance == 2_000_000 + + +@pytest.mark.asyncio +async def test_endpoint_zero_balance_with_open_claim_returns_409( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """A prior refund debited the balance and is still settling: the retry must + report refund_in_progress (409), not "no balance to refund" (400).""" + await _open_ambiguous(integration_session, "quote-123") + assert (await _load_key(integration_session)).balance == 0 + + send = AsyncMock() + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()), + patch("routstr.refund.send_to_lnurl", send), + ): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert exc_info.value.status_code == 409 + assert isinstance(exc_info.value.detail, dict) + assert exc_info.value.detail["error"]["code"] == "refund_in_progress" + send.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_cashu_token_survives_failed_ledger_write( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """The token is issued once the mint signs it; a failed cashu_transactions + insert must neither fail the request nor release the balance, and a retry + must replay the token from the claim row.""" + await _seed_key(integration_session) + with ( + patch("routstr.refund.send_token", AsyncMock(return_value="cashuAtoken")), + patch("routstr.refund.token_mint_url", lambda token, mint: mint), + patch( + "routstr.refund.store_cashu_transaction", + AsyncMock(side_effect=RuntimeError("db down")), + ), + ): + first = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(first, dict) + assert (first["token"], first["status"]) == ("cashuAtoken", "paid") + assert (await _load_key(integration_session)).balance == 0 + + second = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(second, dict) + assert second["refund_id"] == first["refund_id"] + assert second["token"] == "cashuAtoken" + + +async def _topup(session: AsyncSession, amount: int = BALANCE_MSATS) -> None: + key = await _load_key(session) + key.balance = amount + session.add(key) + await session.commit() + + +async def _refund_cashu( + session: AsyncSession, token: str, *, ledger: bool +) -> dict[str, str]: + store = ( + AsyncMock(side_effect=RuntimeError("db down")) + if not ledger + else store_cashu_transaction_with_retry + ) + with ( + patch("routstr.refund.send_token", AsyncMock(return_value=token)), + patch("routstr.refund.token_mint_url", lambda t, mint: mint), + patch("routstr.refund.store_cashu_transaction", store), + ): + body = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=session, + ) + assert isinstance(body, dict) + return body + + +@pytest.mark.asyncio +async def test_replay_prefers_the_token_of_the_newest_paid_claim( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """An older ledger row must not answer for a newer claim whose write failed.""" + await _seed_key(integration_session) + await _refund_cashu(integration_session, "cashuAold", ledger=True) + await _topup(integration_session) + newest = await _refund_cashu(integration_session, "cashuAnew", ledger=False) + + replay = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(replay, dict) + assert replay["token"] == "cashuAnew" + assert replay["refund_id"] == newest["refund_id"] + + +@pytest.mark.asyncio +async def test_swept_older_ledger_row_does_not_reject_the_newest_claim( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + await _seed_key(integration_session) + await _refund_cashu(integration_session, "cashuAold", ledger=True) + await _topup(integration_session) + await _refund_cashu(integration_session, "cashuAnew", ledger=False) + + result = await integration_session.exec( + select(CashuTransaction).where(CashuTransaction.token == "cashuAold") + ) + old_tx = result.one() + old_tx.swept = True + integration_session.add(old_tx) + await integration_session.commit() + + replay = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(replay, dict) + assert replay["token"] == "cashuAnew" + + +@pytest.mark.asyncio +async def test_open_claim_is_reported_over_an_older_paid_claim( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """A paid claim from a previous cycle must not be replayed as the outcome + of the claim that is still settling.""" + await _seed_key(integration_session, address=ADDRESS) + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()), + patch("routstr.refund.send_to_lnurl", _lnurl_stub()), + ): + paid = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(paid, dict) + await _topup(integration_session) + key = await _load_key(integration_session) + open_claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + await refund.hold(integration_session, open_claim, "quote-open") + + validate = AsyncMock() + with patch("routstr.refund.get_lnurl_data", validate): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + refund_request=RefundRequest(lightning_address=ADDRESS), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + + assert exc_info.value.status_code == 409 + assert isinstance(exc_info.value.detail, dict) + error = exc_info.value.detail["error"] + assert (error["refund_id"], error["status"]) == (open_claim.id, "ambiguous") + # A drained key must not be able to drive outbound destination lookups. + validate.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_stuck_claim_is_reported_instead_of_no_balance( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + await refund._close(integration_session, claim, status="stuck") + await integration_session.commit() + + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert exc_info.value.status_code == 409 + assert isinstance(exc_info.value.detail, dict) + error = exc_info.value.detail["error"] + assert (error["code"], error["refund_id"]) == ("refund_unresolved", claim.id) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("status", "counted"), + [ + ("pending", True), + ("ambiguous", True), + ("stuck", True), + ("paid", False), + ("failed", False), + ], +) +async def test_liability_covers_claims_until_they_resolve( + integration_session: AsyncSession, + patched_db_engine: None, + status: str, + counted: bool, +) -> None: + """Money in flight is still owed to the customer; an owner payout that read + only key balances could spend its backing.""" + key = await _seed_key(integration_session) + before = await total_user_liability(integration_session) + assert before == BALANCE_MSATS + + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + assert await total_user_liability(integration_session) == BALANCE_MSATS + + await refund._close(integration_session, claim, status=status) + await integration_session.commit() + expected = BALANCE_MSATS if counted else 0 + assert await total_user_liability(integration_session) == expected + + +@pytest.mark.asyncio +async def test_unreachable_destination_is_a_client_error() -> None: + with patch( + "routstr.refund.get_lnurl_data", + AsyncMock(side_effect=httpx.ConnectError("All connection attempts failed")), + ): + with pytest.raises(HTTPException) as exc_info: + await refund.validate_lightning_destination(ADDRESS) + assert exc_info.value.status_code == 400 diff --git a/tests/integration/test_refund_payout_recovery.py b/tests/integration/test_refund_payout_recovery.py new file mode 100644 index 00000000..fb4d37bc --- /dev/null +++ b/tests/integration/test_refund_payout_recovery.py @@ -0,0 +1,368 @@ +"""Refund payouts that finish after the claim row stopped cooperating. + +Each test pins one guarantee of the claim table that the happy path cannot +exercise: the payout side has authoritative knowledge of what the mint did, +and the claim row must end up agreeing with it. +""" + +import time +from typing import Any +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi import HTTPException +from sqlalchemy import event +from sqlalchemy.exc import OperationalError + +from routstr import refund +from routstr.balance import RefundRequest, refund_wallet_endpoint +from routstr.core.db import ApiKey, AsyncSession, Refund, total_user_liability +from routstr.payment.lnurl import MeltUnpaidError + +KEY_HASH = "refundrecoverykey" +ADDRESS = "user@ln.example.com" +BALANCE_MSATS = 5_000_000 + + +async def _seed_key(session: AsyncSession, *, address: str | None = None) -> ApiKey: + key = ApiKey(hashed_key=KEY_HASH) + key.balance = BALANCE_MSATS + key.reserved_balance = 0 + key.refund_currency = "sat" + key.refund_address = address + key.total_spent = 0 + key.total_requests = 0 + session.add(key) + await session.commit() + await session.refresh(key) + return key + + +async def _load_key(session: AsyncSession) -> ApiKey: + key = await session.get(ApiKey, KEY_HASH) + assert key is not None + await session.refresh(key) + return key + + +async def _load_refund(session: AsyncSession, refund_id: str) -> Refund: + row = await session.get(Refund, refund_id) + assert row is not None + await session.refresh(row) + return row + + +def _cashu_payout(token: str, send: Any | None = None) -> Any: + return ( + patch("routstr.refund.send_token", send or AsyncMock(return_value=token)), + patch("routstr.refund.token_mint_url", lambda t, mint: mint), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + ) + + +# --- cashu token issued after the reconciler gave up ----------------------- + + +@pytest.mark.asyncio +async def test_cashu_token_issued_after_reconciler_marked_claim_stuck_settles_it( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """A token in the customer's hands must leave its claim ``paid``: a row + left ``stuck`` keeps the amount in liability forever and tells the operator + to reconcile a payout that already happened.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + + async def slow_send_token(amount: int, unit: str, mint_url: str) -> str: + # The lease lapses while the mint is still working. + row = await _load_refund(integration_session, claim.id) + row.claimed_at = (row.claimed_at or 0) - 10_000 + integration_session.add(row) + await integration_session.commit() + await refund.reconcile_once() + assert (await _load_refund(integration_session, claim.id)).status == "stuck" + return "cashuAlate" + + send, mint, store = _cashu_payout("cashuAlate", slow_send_token) + with send, mint, store: + body = await refund.execute(integration_session, claim) + + assert (body["token"], body["status"]) == ("cashuAlate", "paid") + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.token, row.claimed_at) == ("paid", "cashuAlate", None) + assert await total_user_liability(integration_session) == 0 + + +@pytest.mark.asyncio +async def test_cashu_payout_renews_lease_before_asking_the_mint( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """The reconciler leaves a claim alone while its lease is fresh, so the + payout renews it right before the slow step.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + row = await _load_refund(integration_session, claim.id) + row.claimed_at = (row.claimed_at or 0) - 10_000 + integration_session.add(row) + await integration_session.commit() + + async def send_token(amount: int, unit: str, mint_url: str) -> str: + await refund.reconcile_once() + return "cashuAfresh" + + send, mint, store = _cashu_payout("cashuAfresh", send_token) + with send, mint, store, patch("routstr.refund.logger") as log: + await refund.execute(integration_session, claim) + + log.critical.assert_not_called() + assert (await _load_refund(integration_session, claim.id)).status == "paid" + + +# --- cashu token issued, claim write failed -------------------------------- + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure_point", ["execute", "autoflush", "commit"]) +async def test_cashu_claim_write_failure_after_token_creation_withholds_balance( + integration_session: AsyncSession, + patched_db_engine: None, + integration_engine: Any, + failure_point: str, +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + claim_id = claim.id + fired = False + armed = False + + def fail_once(*args: Any) -> None: + nonlocal fired + statement = args[2] if failure_point != "commit" else "COMMIT" + if ( + armed + and not fired + and (statement.startswith("UPDATE refunds") or statement == "COMMIT") + ): + fired = True + raise OperationalError(statement, {}, Exception("database is locked")) + + async def send_token(amount: int, unit: str, mint_url: str) -> str: + nonlocal armed + if failure_point == "autoflush": + # A pending ORM write makes SQLAlchemy invalidate the transaction + # and expire attached objects when the actual SQL execution fails. + claim.updated_at -= 1 + armed = True + return "cashuAstranded" + + event_name = "commit" if failure_point == "commit" else "before_cursor_execute" + event.listen(integration_engine.sync_engine, event_name, fail_once) + send, _, store = _cashu_payout("cashuAstranded", send_token) + try: + with ( + send, + store, + patch( + "routstr.refund.token_mint_url", return_value="https://fallback.mint" + ), + ): + with pytest.raises(HTTPException) as exc_info: + await refund.execute(integration_session, claim) + finally: + event.remove(integration_engine.sync_engine, event_name, fail_once) + + assert fired + assert exc_info.value.status_code == 502 + row = await _load_refund(integration_session, claim_id) + assert (row.status, row.token, row.mint_url) == ( + "ambiguous", + "cashuAstranded", + "https://fallback.mint", + ) + assert (await _load_key(integration_session)).balance == 0 + assert await total_user_liability(integration_session) == BALANCE_MSATS + + await refund.reconcile_once() + row = await _load_refund(integration_session, claim_id) + assert (row.status, row.token) == ("paid", "cashuAstranded") + assert await total_user_liability(integration_session) == 0 + with patch("routstr.refund.send_token", AsyncMock()) as send_again: + replay = await refund_wallet_endpoint( + refund_request=RefundRequest(), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(replay, dict) + assert replay["token"] == "cashuAstranded" + send_again.assert_not_awaited() + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_reconciler_settles_held_cashu_claim_that_carries_a_token( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """A held cashu claim whose row carries the token is a completed payout; + the reconciler closes it as paid instead of escalating it to stuck.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + await refund.hold(integration_session, claim, None, token="cashuAheld") + row = await _load_refund(integration_session, claim.id) + row.claimed_at = None + row.updated_at -= 10_000 + integration_session.add(row) + await integration_session.commit() + + with patch("routstr.refund.logger") as log: + await refund.reconcile_once() + + log.critical.assert_not_called() + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.token) == ("paid", "cashuAheld") + assert await total_user_liability(integration_session) == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("hold_before_lease", [True, False]) +async def test_stale_cashu_reconciliation_preserves_newly_held_token( + integration_session: AsyncSession, + patched_db_engine: None, + integration_engine: Any, + hold_before_lease: bool, +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + claim_id = claim.id + now = int(time.time()) + cutoff = now - refund.settings.refund_claim_timeout_seconds + claim.claimed_at = cutoff - 1 + await integration_session.commit() + async with AsyncSession(integration_engine, expire_on_commit=False) as session: + stale = await session.get(Refund, claim_id) + assert stale is not None and stale.token is None + + if hold_before_lease: + await refund.hold(integration_session, claim, None, token="cashuAheld") + assert await refund._lease(claim_id, now, cutoff) + real_close = refund._close + + async def hold_before_close(session: AsyncSession, row: Refund, **kw: Any) -> bool: + if not hold_before_lease and kw.get("status") == "stuck": + # Token arrives at the last moment, even after a potential reload. + await real_close( + integration_session, claim, status="ambiguous", token="cashuAheld" + ) + await integration_session.commit() + return await real_close(session, row, **kw) + + with patch.object(refund, "_close", hold_before_close): + await refund._reconcile(stale, now) + + row = await _load_refund(integration_session, claim_id) + assert (row.status, row.token) == ("paid", "cashuAheld") + assert (await _load_key(integration_session)).balance == 0 + assert await total_user_liability(integration_session) == 0 + with patch("routstr.refund.send_token", AsyncMock()) as send_again: + replay = await refund_wallet_endpoint( + refund_request=RefundRequest(), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(replay, dict) + assert replay["token"] == "cashuAheld" + send_again.assert_not_awaited() + + +# --- mint proved the melt unpaid ------------------------------------------- + + +@pytest.mark.asyncio +async def test_mint_confirmed_unpaid_melt_restores_balance_immediately( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """The mint answering ``unpaid`` to the melt itself is proof no payment + happened, so the customer gets the balance back now, not after the + reconciler timeout.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + + async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int: + await on_melt_quote("quote-unpaid", claim.mint_url) + raise MeltUnpaidError("Cashu mint confirmed that the melt was unpaid") + + with patch("routstr.refund.send_to_lnurl", send): + with pytest.raises(HTTPException) as exc_info: + await refund.execute(integration_session, claim) + + assert exc_info.value.status_code == 503 + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.quote_id) == ("failed", "quote-unpaid") + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +# --- response reflects the persisted claim --------------------------------- + + +@pytest.mark.asyncio +async def test_response_status_reflects_persisted_claim( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """When ``settle`` closes nothing the response must not invent ``paid``.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + + async def send_token(amount: int, unit: str, mint_url: str) -> str: + # Somebody closed the row as failed while the mint was working. + await refund._close(integration_session, claim, status="failed") + await integration_session.commit() + return "cashuAorphan" + + send, mint, store = _cashu_payout("cashuAorphan", send_token) + with send, mint, store: + body = await refund.execute(integration_session, claim) + + assert body["status"] == "failed" + assert (await _load_refund(integration_session, claim.id)).status == "failed" + + +@pytest.mark.asyncio +async def test_stored_refund_address_is_validated_before_debit( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + """A bad address stored on the key is a client error, not a payout failure.""" + await _seed_key(integration_session, address="nobody@invalid.example") + send = AsyncMock() + with ( + patch( + "routstr.refund.get_lnurl_data", + AsyncMock(side_effect=refund.LNURLError("no such user")), + ), + patch("routstr.refund.send_to_lnurl", send), + ): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + refund_request=RefundRequest(), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + + assert exc_info.value.status_code == 400 + send.assert_not_awaited() + assert (await _load_key(integration_session)).balance == BALANCE_MSATS diff --git a/tests/integration/test_reserved_balance_negative.py b/tests/integration/test_reserved_balance_negative.py index 8ffebc86..49e290cb 100644 --- a/tests/integration/test_reserved_balance_negative.py +++ b/tests/integration/test_reserved_balance_negative.py @@ -133,14 +133,11 @@ async def test_reserved_balance_with_successful_requests( @pytest.mark.asyncio -async def test_revert_with_zero_reserved_balance_is_noop( +async def test_revert_with_zero_reserved_balance_repairs_terminally( integration_session: AsyncSession, ) -> None: - """Test that revert_pay_for_request is a no-op when reserved_balance is 0. - - Previously this would drive reserved_balance negative. With the floor guard, - it should return False and leave reserved_balance at 0. - """ + """Reverting after the aggregate was already zeroed must not drive it + negative: the corrupt durable reservation is released without subtraction.""" from routstr.auth import pay_for_request, revert_pay_for_request unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}" @@ -153,21 +150,24 @@ async def test_revert_with_zero_reserved_balance_is_noop( await integration_session.commit() await pay_for_request(test_key, 100, integration_session) test_key.reserved_balance = 0 + test_key.total_requests = 0 integration_session.add(test_key) await integration_session.commit() - # A stale cleanup already released the aggregate reservation. + # A stale cleanup already released the aggregate reservation. The revert + # terminalizes the durable row (repair) without driving the aggregate + # negative. result = await revert_pay_for_request(test_key, integration_session, 100) - await integration_session.refresh(test_key) + integration_session.expunge_all() + updated = await integration_session.get(ApiKey, unique_key) + assert updated is not None - assert result is False, "Revert should return False when reservation already released" - assert test_key.reserved_balance == 0, ( - f"Reserved balance should remain 0, got: {test_key.reserved_balance}" - ) - assert test_key.total_requests == 1, ( - f"Total requests should remain 1, got: {test_key.total_requests}" + assert result is True, "Revert must terminalize the corrupt reservation" + assert updated.reserved_balance == 0, ( + f"Reserved balance should remain 0, got: {updated.reserved_balance}" ) + assert updated.total_requests == 0 @pytest.mark.asyncio @@ -203,10 +203,11 @@ async def test_revert_with_sufficient_reserved_balance_succeeds( @pytest.mark.asyncio -async def test_revert_partial_reserved_balance_is_noop( +async def test_revert_partial_reserved_balance_repairs_terminally( integration_session: AsyncSession, ) -> None: - """Test that reverting more than the current reserved_balance is a no-op.""" + """Reverting more than the aggregate holds must not clamp or go negative: + the corrupt durable reservation is released without subtraction.""" from routstr.auth import pay_for_request, revert_pay_for_request unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}" @@ -223,18 +224,19 @@ async def test_revert_partial_reserved_balance_is_noop( integration_session.add(test_key) await integration_session.commit() - # Try to revert 500 when only 50 is reserved — should be no-op + # Reverting 500 when only 50 is reserved cannot subtract; the corrupt + # reservation is terminalized and the aggregate left untouched. result = await revert_pay_for_request(test_key, integration_session, 500) - await integration_session.refresh(test_key) + integration_session.expunge_all() + updated = await integration_session.get(ApiKey, unique_key) + assert updated is not None - assert result is False, "Revert should fail when cost > reserved_balance" - assert test_key.reserved_balance == 50, ( - f"Reserved balance should stay at 50, got: {test_key.reserved_balance}" - ) - assert test_key.total_requests == 1, ( - f"Total requests should stay at 1, got: {test_key.total_requests}" + assert result is True, "Revert must terminalize the corrupt reservation" + assert updated.reserved_balance == 50, ( + f"Reserved balance should stay at 50, got: {updated.reserved_balance}" ) + assert updated.total_requests == 0 @pytest.mark.asyncio @@ -265,9 +267,7 @@ async def test_double_revert_prevented( snapshot = await get_reservation_snapshot(test_key, integration_session) # First revert — should succeed - result1 = await revert_pay_for_request( - test_key, integration_session, 500, snapshot - ) + result1 = await revert_pay_for_request(test_key, integration_session, 500, snapshot) await integration_session.refresh(test_key) assert result1 is True @@ -275,9 +275,7 @@ async def test_double_revert_prevented( assert test_key.total_requests == 4 # Second revert of the same amount — should be no-op - result2 = await revert_pay_for_request( - test_key, integration_session, 500, snapshot - ) + result2 = await revert_pay_for_request(test_key, integration_session, 500, snapshot) await integration_session.refresh(test_key) assert result2 is False, "Second revert should be a no-op" @@ -319,9 +317,7 @@ async def test_sequential_reverts_never_go_negative( # Run 5 sequential reverts for the same 500 reservation results = [] for _ in range(5): - r = await revert_pay_for_request( - test_key, integration_session, 500, snapshot - ) + r = await revert_pay_for_request(test_key, integration_session, 500, snapshot) results.append(r) await integration_session.refresh(test_key) @@ -337,63 +333,3 @@ async def test_sequential_reverts_never_go_negative( assert test_key.reserved_balance >= 0, ( f"Reserved balance went negative: {test_key.reserved_balance}" ) - - -@pytest.mark.asyncio -async def test_child_key_revert_floor_guard( - integration_session: AsyncSession, -) -> None: - """Test that child key reserved_balance also has floor guard on revert.""" - from routstr.auth import ( - get_reservation_snapshot, - pay_for_request, - revert_pay_for_request, - ) - - parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}" - child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}" - - parent_key = ApiKey( - hashed_key=parent_key_hash, - balance=10000, - reserved_balance=0, - total_requests=2, - ) - child_key = ApiKey( - hashed_key=child_key_hash, - balance=0, - reserved_balance=0, - total_requests=2, - parent_key_hash=parent_key_hash, - ) - integration_session.add(parent_key) - integration_session.add(child_key) - await integration_session.commit() - await pay_for_request(child_key, 500, integration_session) - snapshot = await get_reservation_snapshot(child_key, integration_session) - - # First revert succeeds - result1 = await revert_pay_for_request( - child_key, integration_session, 500, snapshot - ) - await integration_session.refresh(parent_key) - await integration_session.refresh(child_key) - - assert result1 is True - assert parent_key.reserved_balance == 0 - assert child_key.reserved_balance == 0 - - # Second revert is a no-op for both parent and child - result2 = await revert_pay_for_request( - child_key, integration_session, 500, snapshot - ) - await integration_session.refresh(parent_key) - await integration_session.refresh(child_key) - - assert result2 is False - assert parent_key.reserved_balance == 0, ( - f"Parent reserved_balance should stay 0, got: {parent_key.reserved_balance}" - ) - assert child_key.reserved_balance == 0, ( - f"Child reserved_balance should stay 0, got: {child_key.reserved_balance}" - ) diff --git a/tests/integration/test_routstr_auto_topup_claim.py b/tests/integration/test_routstr_auto_topup_claim.py new file mode 100644 index 00000000..dc246a63 --- /dev/null +++ b/tests/integration/test_routstr_auto_topup_claim.py @@ -0,0 +1,380 @@ +"""Real-database tests for the Routstr-to-Routstr auto top-up spend bound. + +The bound has to survive a process restart and concurrent workers, so these +run against actual SQL instead of mocked sessions: an in-memory counter would +pass a mocked test and still let a non-crediting peer drain the wallet. +""" + +import importlib +import time +from contextlib import ExitStack +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from sqlmodel import select + +from routstr.core.db import CashuTransaction, UpstreamProviderRow, create_session +from routstr.upstream import auto_topup as auto_topup_module +from routstr.upstream.auto_topup import ( + ROUTSTR_MAX_DAILY_TOPUP_SATS, + ROUTSTR_MAX_TOPUP_FAILURES, + ROUTSTR_PHASE_BACKOFF, + ROUTSTR_PHASE_HALTED, + ROUTSTR_PHASE_SENT, + _check_and_topup, + _claim_routstr_topup, + _parse_routstr_request_id, + _persist_routstr_token_and_mark_sent, + _routstr_spent_last_24h_sats, + _routstr_state_id_for_provider, + get_routstr_auto_topup_state, + release_routstr_auto_topup_state, +) + +pytestmark = pytest.mark.asyncio + +TOPUP_SATS = 50 + + +async def _seed_provider(provider_id: int = 1) -> UpstreamProviderRow: + import json + + row = UpstreamProviderRow( + id=provider_id, + slug=f"peer-{provider_id}", + provider_type="routstr", + base_url="https://peer.test", + api_key="secret", + enabled=True, + provider_settings=json.dumps( + { + "auto_topup": True, + "topup_threshold": 1, + "topup_amount_limit": TOPUP_SATS, + "topup_mint_url": "https://mint.test", + } + ), + ) + async with create_session() as session: + session.add(row) + await session.commit() + await session.refresh(row) + return row + + +def _peer(balance: float, *, topup: object = None) -> MagicMock: + provider = MagicMock() + provider.get_balance = AsyncMock(return_value=balance) + provider.topup = AsyncMock(return_value=topup or {"balance": balance}) + return provider + + +def _patch_wallet(module: Any, peer: MagicMock, token: str) -> ExitStack: + stack = ExitStack() + stack.enter_context( + patch.object(module.RoutstrUpstreamProvider, "from_db_row", return_value=peer) + ) + stack.enter_context( + patch.object( + module, "send_token_from_owner_locked", AsyncMock(return_value=token) + ) + ) + stack.enter_context( + patch.object(module, "token_mint_url", return_value="https://mint.test") + ) + return stack + + +async def _claim_state(provider_id: int = 1) -> CashuTransaction | None: + async with create_session() as session: + return await session.get( + CashuTransaction, _routstr_state_id_for_provider(provider_id) + ) + + +async def _sent_tokens() -> list[CashuTransaction]: + async with create_session() as session: + return list( + ( + await session.exec( + select(CashuTransaction).where( + CashuTransaction.source == "auto_topup" + ) + ) + ).all() + ) + + +async def test_second_worker_cannot_claim_while_the_first_holds_one( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + assert await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) is not None + assert await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) is None + + +async def test_token_and_sent_claim_roll_back_together_on_commit_failure( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + operation_id = await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) + assert operation_id is not None + + with patch( + "sqlmodel.ext.asyncio.session.AsyncSession.commit", + new=AsyncMock(side_effect=RuntimeError("commit failed")), + ): + with pytest.raises(RuntimeError, match="commit failed"): + await _persist_routstr_token_and_mark_sent( + row, + operation_id, + expected_sats=TOPUP_SATS, + token="cashu-token-atomic", + amount=TOPUP_SATS, + mint_url="https://mint.test", + ) + + claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr] + assert claim is not None and claim.phase != ROUTSTR_PHASE_SENT + assert await _sent_tokens() == [] + + +async def test_token_is_persisted_before_it_reaches_the_peer( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + seen: list[CashuTransaction] = [] + + async def _record_then_accept(token: str) -> dict: + seen.extend(await _sent_tokens()) + return {"balance": TOPUP_SATS} + + peer = _peer(0.0) + peer.topup = AsyncMock(side_effect=_record_then_accept) + + with _patch_wallet(auto_topup_module, peer, "cashu-token-1"): + await _check_and_topup(row) + + assert [tx.token for tx in seen] == ["cashu-token-1"] + assert seen[0].collected is False + assert (await _sent_tokens())[0].collected is True + + +async def test_untracked_token_is_returned_and_never_sent( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + peer = _peer(0.0) + + with ( + _patch_wallet(auto_topup_module, peer, "cashu-token-1"), + patch.object( + auto_topup_module, + "_persist_routstr_token_and_mark_sent", + AsyncMock(side_effect=RuntimeError("database unavailable")), + ), + patch.object( + auto_topup_module, "release_token_reservation", AsyncMock() + ) as reclaim, + ): + await _check_and_topup(row) + + reclaim.assert_awaited_once_with("cashu-token-1") + peer.topup.assert_not_awaited() + # Nothing left the wallet, so the slot must be free again immediately. + state = await _claim_state() + assert state is not None and state.swept is True + + +async def test_failed_topup_keeps_token_uncollected_and_suppresses_new_claims( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + peer = _peer(0.0, topup={"error": "rejected"}) + + with _patch_wallet(auto_topup_module, peer, "cashu-token-1"): + await _check_and_topup(row) + + tokens = await _sent_tokens() + assert len(tokens) == 1 + assert tokens[0].collected is False + + claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr] + assert claim is not None and claim.phase == ROUTSTR_PHASE_SENT + + peer.topup.reset_mock() + with _patch_wallet(auto_topup_module, peer, "cashu-token-2"): + await _check_and_topup(row) + + peer.topup.assert_not_awaited() + assert len(await _sent_tokens()) == 1 + + +async def test_claim_blocks_a_restarted_process(patched_db_engine: Any) -> None: + row = await _seed_provider() + peer = _peer(0.0, topup={"error": "rejected"}) + + with _patch_wallet(auto_topup_module, peer, "cashu-token-1"): + await _check_and_topup(row) + + # A fresh module drops every process-local variable; only a durable row + # can still stop the next payment. + reloaded = importlib.reload(auto_topup_module) + try: + peer.topup.reset_mock() + with _patch_wallet(reloaded, peer, "cashu-token-2"): + await reloaded._check_and_topup(row) + peer.topup.assert_not_awaited() + finally: + importlib.reload(auto_topup_module) + + assert len(await _sent_tokens()) == 1 + + +async def _uncredited_attempt( + row: UpstreamProviderRow, peer: MagicMock, token: str +) -> None: + """One full payment attempt against a peer that never credits it.""" + with _patch_wallet(auto_topup_module, peer, token): + await _check_and_topup(row) + await _expire_claim() + with _patch_wallet(auto_topup_module, peer, f"{token}-retry"): + await _check_and_topup(row) + await _expire_claim() + + +async def test_non_crediting_peer_is_halted_after_repeated_failures( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + peer = _peer(0.0) + + for attempt in range(ROUTSTR_MAX_TOPUP_FAILURES): + await _uncredited_attempt(row, peer, f"cashu-token-{attempt}") + + claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr] + assert claim is not None and claim.phase == ROUTSTR_PHASE_HALTED + + peer.topup.reset_mock() + with _patch_wallet(auto_topup_module, peer, "cashu-token-after-halt"): + await _check_and_topup(row) + peer.topup.assert_not_awaited() + assert len(await _sent_tokens()) == ROUTSTR_MAX_TOPUP_FAILURES + + +async def _expire_claim(provider_id: int = 1) -> None: + """Age the claim's deadline so the reconciler treats it as timed out.""" + async with create_session() as session: + state = await session.get( + CashuTransaction, _routstr_state_id_for_provider(provider_id) + ) + assert state is not None and state.request_id is not None + claim = _parse_routstr_request_id(state.request_id) + assert claim is not None + state.request_id = auto_topup_module._routstr_request_id( + claim.operation_id, + int(time.time()) - 1, + claim.phase, + claim.expected_sats, + claim.failures, + ) + session.add(state) + await session.commit() + + +async def test_crediting_peer_releases_the_claim_for_a_later_topup( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + peer = _peer(0.0) + + with _patch_wallet(auto_topup_module, peer, "cashu-token-1"): + await _check_and_topup(row) + + tokens = await _sent_tokens() + assert len(tokens) == 1 and tokens[0].collected is True + + peer.get_balance = AsyncMock(return_value=float(TOPUP_SATS)) + with _patch_wallet(auto_topup_module, peer, "cashu-token-2"): + await _check_and_topup(row) + + state = await _claim_state() + assert state is not None and state.collected is True + + +async def test_rolling_budget_refuses_a_topup_that_would_exceed_it( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + async with create_session() as session: + session.add( + CashuTransaction( + id="prior-spend", + token="cashu-prior", + amount=ROUTSTR_MAX_DAILY_TOPUP_SATS, + unit="sat", + type="out", + source="auto_topup", + ) + ) + await session.commit() + + assert await _routstr_spent_last_24h_sats() == ROUTSTR_MAX_DAILY_TOPUP_SATS + + peer = _peer(0.0) + with _patch_wallet(auto_topup_module, peer, "cashu-token-1"): + await _check_and_topup(row) + + peer.topup.assert_not_awaited() + assert await _claim_state() is None + + +async def test_admin_release_is_fenced_on_the_state_it_reviewed( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + peer = _peer(0.0, topup={"error": "rejected"}) + with _patch_wallet(auto_topup_module, peer, "cashu-token-1"): + await _check_and_topup(row) + + state = await get_routstr_auto_topup_state(1) + assert state["active"] is True + assert state["phase"] == ROUTSTR_PHASE_SENT + + stale = await release_routstr_auto_topup_state( + 1, state_token="routstr:other:0:sent:0:0" + ) + assert stale.released is False and stale.reason == "stale_state" + + released = await release_routstr_auto_topup_state( + 1, state_token=str(state["state_token"]) + ) + assert released.released is True + + peer.topup.reset_mock() + with _patch_wallet(auto_topup_module, peer, "cashu-token-2"): + await _check_and_topup(row) + peer.topup.assert_awaited_once() + + +async def test_backoff_suppresses_retries_until_its_deadline( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + peer = _peer(0.0) + + with _patch_wallet(auto_topup_module, peer, "cashu-token-1"): + await _check_and_topup(row) + await _expire_claim() + + # First reconciliation after the lease: the peer never credited, so the + # claim moves to backoff rather than paying again immediately. + peer.topup.reset_mock() + with _patch_wallet(auto_topup_module, peer, "cashu-token-2"): + await _check_and_topup(row) + peer.topup.assert_not_awaited() + + claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr] + assert claim is not None and claim.phase == ROUTSTR_PHASE_BACKOFF + assert claim.deadline > int(time.time()) diff --git a/tests/integration/test_routstr_auto_topup_threshold_units.py b/tests/integration/test_routstr_auto_topup_threshold_units.py new file mode 100644 index 00000000..a680b31d --- /dev/null +++ b/tests/integration/test_routstr_auto_topup_threshold_units.py @@ -0,0 +1,164 @@ +"""The Routstr top-up threshold has to state its unit. + +``topup_amount_limit`` was already plain sats — it goes straight to +``send_token(amount, "sat", ...)`` and is added to the balance for logging — but +its sibling ``topup_threshold`` was compared as ``balance >= threshold * 1000``. +Two keys from one settings blob, two different units, and nothing naming either. +An operator asking for "top up below 1000 sats" got one below 1,000,000. + +``topup_threshold_sats`` says what it means. Legacy values keep their effective +trigger point: reinterpreting them as sats would silently drop the trigger a +thousandfold and let a provider run dry. +""" + +import json +from typing import Any +from unittest.mock import patch + +import pytest + +from routstr.core.db import UpstreamProviderRow, create_session +from routstr.upstream import auto_topup as auto_topup_module +from routstr.upstream.auto_topup import ( + _check_and_topup, + validate_routstr_auto_topup_settings, +) + +from .test_routstr_auto_topup_claim import _patch_wallet, _peer, _sent_tokens + +# No module-level asyncio mark: the settings cases are sync, and the suite +# already runs asyncio in auto mode. + +TOPUP_SATS = 50 + + +@pytest.fixture(autouse=True) +def _forget_legacy_hints() -> Any: + # The hint fires once per provider for the life of the process. Reach it + # through the module: a sibling test reloads auto_topup, which rebinds the + # set, so a name imported here would go on clearing the old one. + auto_topup_module._legacy_threshold_hinted.clear() + yield + auto_topup_module._legacy_threshold_hinted.clear() + + +async def _seed(**topup_settings: Any) -> UpstreamProviderRow: + row = UpstreamProviderRow( + id=1, + slug="peer-1", + provider_type="routstr", + base_url="https://peer.test", + api_key="secret", + enabled=True, + provider_settings=json.dumps( + { + "auto_topup": True, + "topup_amount_limit": TOPUP_SATS, + "topup_mint_url": "https://mint.test", + **topup_settings, + } + ), + ) + async with create_session() as session: + session.add(row) + await session.commit() + await session.refresh(row) + return row + + +async def _topped_up(row: UpstreamProviderRow, balance: float) -> bool: + peer = _peer(balance) + with _patch_wallet(auto_topup_module, peer, "cashu-token-1"): + await _check_and_topup(row) + return bool(await _sent_tokens()) + + +async def test_explicit_sats_threshold_is_compared_against_a_sats_balance( + patched_db_engine: Any, +) -> None: + row = await _seed(topup_threshold_sats=1000) + assert await _topped_up(row, 1500.0) is False + + +async def test_explicit_sats_threshold_tops_up_below_the_stated_amount( + patched_db_engine: Any, +) -> None: + row = await _seed(topup_threshold_sats=1000) + assert await _topped_up(row, 500.0) is True + + +@pytest.mark.parametrize( + ("balance", "expected"), + [(1500.0, False), (500.0, True)], +) +async def test_legacy_threshold_keeps_its_effective_trigger_point( + patched_db_engine: Any, balance: float, expected: bool +) -> None: + # 1 x 1000 == the 1000 sats the old comparison actually used. Reading it as + # 1 sat instead would leave the peer to run dry. + row = await _seed(topup_threshold=1) + assert await _topped_up(row, balance) is expected + + +async def test_explicit_sats_threshold_overrides_the_legacy_key( + patched_db_engine: Any, +) -> None: + row = await _seed(topup_threshold=1, topup_threshold_sats=100) + assert await _topped_up(row, 500.0) is False + + +async def test_legacy_threshold_reports_the_sats_value_to_migrate_to( + patched_db_engine: Any, +) -> None: + row = await _seed(topup_threshold=1) + # The app logger does not propagate to root, so caplog would see nothing. + with patch.object(auto_topup_module, "logger") as log: + await _topped_up(row, 5000.0) + hints = [ + c for c in log.warning.call_args_list if "topup_threshold_sats" in c.args[0] + ] + assert len(hints) == 1 + assert hints[0].kwargs["extra"]["threshold_sats"] == 1000.0 + + # The scheduler runs every minute; the hint must not run with it. + await _topped_up(row, 5000.0) + assert ( + sum("topup_threshold_sats" in c.args[0] for c in log.warning.call_args_list) + == 1 + ) + + +def test_settings_accept_the_sats_threshold_on_its_own() -> None: + assert ( + validate_routstr_auto_topup_settings( + { + "auto_topup": True, + "topup_threshold_sats": 1000, + "topup_amount_limit": TOPUP_SATS, + "topup_mint_url": "https://mint.test", + } + ) + is None + ) + + +@pytest.mark.parametrize("value", [0, -1, True, float("inf"), "1000", None]) +def test_settings_reject_an_unusable_sats_threshold(value: object) -> None: + assert validate_routstr_auto_topup_settings( + { + "auto_topup": True, + "topup_threshold_sats": value, + "topup_amount_limit": TOPUP_SATS, + "topup_mint_url": "https://mint.test", + } + ) + + +def test_settings_require_one_of_the_threshold_keys() -> None: + assert validate_routstr_auto_topup_settings( + { + "auto_topup": True, + "topup_amount_limit": TOPUP_SATS, + "topup_mint_url": "https://mint.test", + } + ) diff --git a/tests/integration/test_served_catalog_rate_backstop.py b/tests/integration/test_served_catalog_rate_backstop.py new file mode 100644 index 00000000..4acee442 --- /dev/null +++ b/tests/integration/test_served_catalog_rate_backstop.py @@ -0,0 +1,255 @@ +"""The served catalog is the last guard between a stored row and a charge. + +Stored pricing is JSON written by whatever produced the row — an upstream +import, an operator, a legacy migration, or a foreign writer that never passed +the admin edge. So the read path cannot assume a stored rate is a number: it +must decline to serve a row it cannot bill on, and it must survive a row it +cannot read at all rather than taking the whole catalog down with it. + +The admin listing is deliberately exempt: it includes disabled models and is the +one view that still shows the operator the row that needs repair. +""" + +from __future__ import annotations + +import json + +import pytest +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.payment.models import list_models + +_ARCHITECTURE = json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, + } +) + + +async def _make_provider(session: AsyncSession) -> int: + provider = UpstreamProviderRow( + provider_type="generic", + base_url="https://served-upstream.example/v1", + api_key="test-key", + provider_fee=1.0, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + assert provider.id is not None + return provider.id + + +async def _insert_row( + session: AsyncSession, + provider_id: int, + *, + model_id: str, + pricing: dict[str, object], +) -> None: + session.add( + ModelRow( + id=model_id, + name=model_id, + description="d", + created=0, + context_length=8192, + architecture=_ARCHITECTURE, + pricing=json.dumps(pricing), + upstream_provider_id=provider_id, + enabled=True, + forwarded_model_id=model_id, + ) + ) + await session.commit() + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "bad_rate", + [float("nan"), float("inf"), -1.0], + ids=["nan", "inf", "negative"], +) +async def test_served_catalog_excludes_a_malformed_stored_rate( + integration_session: AsyncSession, bad_rate: float +) -> None: + """A stored rate that is not a number must not be advertised. + + Zero is a real price and a free model is servable, but a negative or + non-finite rate is not a price at all: serving it advertises a rate the cost + calculation cannot bill on, so every request falls through to the flat + maximum reservation — or, for a negative rate, bills an amount settlement + credits back to the caller. + """ + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="good", + pricing={"prompt": 1e-06, "completion": 2e-06}, + ) + await _insert_row( + integration_session, + provider_id, + model_id="bad-rate", + pricing={"prompt": bad_rate, "completion": 2e-06}, + ) + + served = {m.id for m in await list_models(integration_session, provider_id)} + + assert served == {"good"} + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_free_stored_price_is_still_served( + integration_session: AsyncSession, +) -> None: + """Zero is a real price. Rejecting malformed rates must not also drop a row + priced at zero, which is a free model and not a broken one.""" + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="free", + pricing={"prompt": 0.0, "completion": 0.0}, + ) + + served = {m.id for m in await list_models(integration_session, provider_id)} + + assert served == {"free"} + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_listing_still_shows_a_malformed_stored_rate( + integration_session: AsyncSession, +) -> None: + """The operator has to be able to see the row that needs fixing. + + The backstop keeps a malformed row out of the *served* catalog. The listing + that includes disabled models is the one view where the row must still + appear, or the operator loses the ability to repair it. + """ + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="bad-rate", + pricing={"prompt": -1.0, "completion": 2e-06}, + ) + + listed = { + m.id + for m in await list_models( + integration_session, provider_id, include_disabled=True + ) + } + + assert listed == {"bad-rate"} + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_one_unreadable_stored_price_does_not_blank_the_catalog( + integration_session: AsyncSession, +) -> None: + """A single unparseable row must cost that row, not every model on the node. + + Stored pricing is JSON written by whatever produced the row, so a + non-numeric rate is reachable from a legacy import or a foreign writer. + Parsing it raises out of the row-to-model conversion, and because the + conversion ran inside the catalog loop the exception took the whole listing + with it — one bad row and the node advertised nothing at all. + """ + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="good", + pricing={"prompt": 1e-06, "completion": 2e-06}, + ) + await _insert_row( + integration_session, + provider_id, + model_id="unreadable", + pricing={"prompt": "not-a-number", "completion": 2e-06}, + ) + + served = {m.id for m in await list_models(integration_session, provider_id)} + + assert served == {"good"} + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize( + "field", + [ + "image", + "web_search", + "internal_reasoning", + "input_cache_read", + "input_cache_write", + ], +) +async def test_served_catalog_excludes_a_malformed_auxiliary_rate( + integration_session: AsyncSession, field: str +) -> None: + """The backstop covers every billable rate, not only the token rates. + + A price whose ``prompt``/``completion`` are sound can still carry a + malformed request, image, search, reasoning or cache rate — the catalog + import filter never inspects those — and the request that hits one is billed + against it just the same. + """ + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="good", + pricing={"prompt": 1e-06, "completion": 2e-06}, + ) + await _insert_row( + integration_session, + provider_id, + model_id="bad-aux", + pricing={"prompt": 1e-06, "completion": 2e-06, field: -1.0}, + ) + + served = {m.id for m in await list_models(integration_session, provider_id)} + + assert served == {"good"} + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_a_negative_request_rate_is_clamped_on_read_and_still_served( + integration_session: AsyncSession, +) -> None: + """``request`` is the one billable rate the row-to-model conversion repairs. + + It clamps a negative stored ``request`` to zero before the price is built, + so the backstop never sees one and the row is served at a zero request rate + — money-safe, and the reason ``request`` is absent from the list of rates + above. Pinned here so that if the clamp goes, this rate joins that list + rather than quietly becoming the one unguarded field. + """ + provider_id = await _make_provider(integration_session) + await _insert_row( + integration_session, + provider_id, + model_id="neg-request", + pricing={"prompt": 1e-06, "completion": 2e-06, "request": -1.0}, + ) + + served = await list_models(integration_session, provider_id) + + assert [m.id for m in served] == ["neg-request"] + assert served[0].pricing.request == 0.0 diff --git a/tests/integration/test_swap_fee_retry.py b/tests/integration/test_swap_fee_retry.py deleted file mode 100644 index 8d4317c0..00000000 --- a/tests/integration/test_swap_fee_retry.py +++ /dev/null @@ -1,210 +0,0 @@ -""" -Integration tests for reactive swap fee retries via the wallet topup endpoint. - -Foreign-mint tokens are swapped to the primary mint using the foreign mint's -melt quote, whose fee_reserve is a non-binding estimate (NUT-05): the mint may -demand more when re-quoting or at melt execution. These tests cover the -endpoint behaviour in those cases: - -1. The mint demands one sat more at melt time than every quote reported - (the mint.cubabitcoin.org incident): the swap retries with a smaller - invoice and the topup succeeds, crediting the recomputed amount. -2. The real melt quote reports a higher fee_reserve than the estimate: the - swap re-quotes from the observed fee and the topup succeeds. -3. The mint escalates its fee demands on every attempt: the retry budget is - exhausted and the endpoint returns 400 with a clear error (never 500), - without ever executing a melt. -""" - -from collections.abc import Callable -from unittest.mock import AsyncMock, Mock, patch - -import pytest -from cashu.core.base import MeltQuoteState -from httpx import AsyncClient, Response - -from routstr.core.settings import settings - -# Captured at collection time, before the integration_app fixture replaces it -# with the testmint stub that bypasses swapping (see conftest.py). -from routstr.wallet import recieve_token as _real_recieve_token - -# Match the authenticated fixture's persisted refund mint: existing-key topups -# are intentionally constrained to that mint for collateral provenance. -PRIMARY_MINT = "http://localhost:3338" - - -def _make_swap_mocks( - token_amount: int, - fee_reserves: list[int], - input_fees: int = 0, - mint_url: str = "http://foreign-mint:3338", -) -> tuple[Mock, Mock, Mock]: - """Return (token, token_wallet, primary_wallet) mocks that act like a mint. - - Mint quotes pass the requested amount through their ``request`` field and - melt quotes echo that amount back, so the mocks stay consistent for - whatever amounts the implementation requests. ``fee_reserves`` supplies the - fee_reserve of each successive melt quote (the first serves the estimation - pass); requesting more quotes than provided fails the test. - """ - mock_token = Mock() - mock_token.mint = mint_url - mock_token.unit = "sat" - mock_token.amount = token_amount - mock_token.keysets = ["keyset1"] - mock_token.proofs = [Mock(amount=token_amount)] - - mock_token_wallet = Mock() - mock_token_wallet.load_mint_keysets = AsyncMock() - mock_token_wallet.activate_keyset = AsyncMock() - mock_token_wallet._expand_short_keyset_ids = AsyncMock() - mock_token_wallet.load_proofs = AsyncMock() - mock_token_wallet.get_fees_for_proofs = Mock(return_value=input_fees) - - mock_primary_wallet = Mock() - mock_primary_wallet.load_mint = AsyncMock() - mock_primary_wallet.load_proofs = AsyncMock() - mock_primary_wallet.available_balance = Mock(amount=0) - mock_primary_wallet.mint = AsyncMock(return_value=Mock()) - - fees = iter(fee_reserves) - - def _next_fee() -> int: - try: - return next(fees) - except StopIteration: - raise AssertionError( - "more melt quotes requested than fee_reserves provided" - ) from None - - mock_primary_wallet.request_mint = AsyncMock( - side_effect=lambda amount: Mock(quote=f"mint_quote_{amount}", request=amount) - ) - mock_token_wallet.melt_quote = AsyncMock( - side_effect=lambda invoice: Mock( - quote=f"melt_quote_{invoice}", amount=invoice, fee_reserve=_next_fee() - ) - ) - mock_token_wallet.melt = AsyncMock( - return_value=Mock(state=MeltQuoteState.paid) - ) - - return mock_token, mock_token_wallet, mock_primary_wallet - - -def _wallet_router(primary_wallet: Mock, token_wallet: Mock) -> Callable[..., Mock]: - """Route get_wallet calls to the primary or foreign wallet mock by URL.""" - - def fake_get_wallet( - mint_url: str, - unit: str = "sat", - load: bool = True, - **kwargs: object, - ) -> Mock: - return primary_wallet if mint_url == PRIMARY_MINT else token_wallet - - return fake_get_wallet - - -async def _post_topup( - client: AsyncClient, - mock_token: Mock, - token_wallet: Mock, - primary_wallet: Mock, -) -> Response: - """POST /v1/wallet/topup with the swap layer mocked at the mint boundary. - - The conftest's testmint stub for recieve_token is swapped back for the - real implementation so the request exercises the actual swap path. - """ - with patch("routstr.wallet.recieve_token", _real_recieve_token): - with patch( - "routstr.wallet.deserialize_token_from_string", return_value=mock_token - ): - with patch( - "routstr.wallet.get_wallet", - side_effect=_wallet_router(primary_wallet, token_wallet), - ): - with patch.object(settings, "primary_mint", PRIMARY_MINT): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch.object(settings, "cashu_mints", [PRIMARY_MINT]): - return await client.post( - "/v1/wallet/topup", - params={"cashu_token": "cashuAtest_foreign_token"}, - ) - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_topup_retries_when_melt_demands_more_than_quoted( - authenticated_client: AsyncClient, -) -> None: - """A 179-sat token where every quote reports fee_reserve=1 but the mint - rejects the first melt demanding 180. The retry shrinks the invoice to 177 - and the topup credits 177 sats (177_000 msats).""" - mock_token, token_wallet, primary_wallet = _make_swap_mocks( - 179, fee_reserves=[1, 1, 1], mint_url="http://mint.cubabitcoin.org" - ) - token_wallet.melt.side_effect = [ - Exception( - "Mint Error: not enough inputs provided for melt. " - "Provided: 179, needed: 180 (Code: 11000)" - ), - Mock(state=MeltQuoteState.paid), - ] - - response = await _post_topup( - authenticated_client, mock_token, token_wallet, primary_wallet - ) - - assert response.status_code == 200 - assert response.json()["msats"] == 177_000 - assert token_wallet.melt.call_count == 2 - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_topup_retries_when_quote_fee_exceeds_estimate( - authenticated_client: AsyncClient, -) -> None: - """A 1000-sat token estimated at fee 20, but the real quote demands 23. - The retry recomputes 1000 - 23 = 977, which fits, and the topup credits - 977 sats (977_000 msats) with a single melt.""" - mock_token, token_wallet, primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[20, 23, 23] - ) - - response = await _post_topup( - authenticated_client, mock_token, token_wallet, primary_wallet - ) - - assert response.status_code == 200 - assert response.json()["msats"] == 977_000 - assert token_wallet.melt.call_count == 1 - - -@pytest.mark.integration -@pytest.mark.asyncio -async def test_topup_returns_422_when_retries_exhausted( - authenticated_client: AsyncClient, -) -> None: - """A mint that escalates fee_reserve on every re-quote (1 → 10 → 25 → 50) - exhausts the retry budget: clean 422 mint_error/too-small taxonomy, melt - never executed.""" - mock_token, token_wallet, primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[1, 10, 25, 50] - ) - - response = await _post_topup( - authenticated_client, mock_token, token_wallet, primary_wallet - ) - - assert response.status_code == 422 - raw_detail = response.json()["detail"] - message = ( - raw_detail["error"]["message"] if isinstance(raw_detail, dict) else raw_detail - ) - assert "too small to cover swap fees" in message - assert token_wallet.melt_quote.call_count == 4 # estimation + 3 attempts - token_wallet.melt.assert_not_called() diff --git a/tests/integration/test_temporary_balances_api.py b/tests/integration/test_temporary_balances_api.py index feee0113..80c0e627 100644 --- a/tests/integration/test_temporary_balances_api.py +++ b/tests/integration/test_temporary_balances_api.py @@ -22,18 +22,18 @@ async def _add_key( hashed_key: str, *, balance: int = 0, + reserved_balance: int = 0, total_spent: int = 0, total_requests: int = 0, created_at: int | None = None, - parent_key_hash: str | None = None, refund_address: str | None = None, ) -> ApiKey: key = ApiKey( hashed_key=hashed_key, balance=balance, + reserved_balance=reserved_balance, total_spent=total_spent, total_requests=total_requests, - parent_key_hash=parent_key_hash, refund_address=refund_address, ) key.created_at = created_at @@ -123,29 +123,18 @@ async def test_temporary_balances_pagination( @pytest.mark.integration @pytest.mark.asyncio -async def test_temporary_balances_totals_exclude_child_balance( +async def test_temporary_balances_totals( integration_client: httpx.AsyncClient, integration_session: AsyncSession, ) -> None: await _add_key( integration_session, - "parent", + "standalone_key", balance=5000, total_spent=100, total_requests=3, created_at=1000, ) - # Child draws from parent's balance, so its balance must NOT be summed, - # but its spent/requests still count. - await _add_key( - integration_session, - "child", - balance=0, - total_spent=200, - total_requests=7, - created_at=1001, - parent_key_hash="parent", - ) response = await integration_client.get( "/admin/api/temporary-balances", headers=_admin_headers() @@ -153,8 +142,8 @@ async def test_temporary_balances_totals_exclude_child_balance( totals = response.json()["totals"] assert totals["total_balance"] == 5000 - assert totals["total_spent"] == 300 - assert totals["total_requests"] == 10 + assert totals["total_spent"] == 100 + assert totals["total_requests"] == 3 @pytest.mark.integration diff --git a/tests/integration/test_topup_untrusted_mint.py b/tests/integration/test_topup_untrusted_mint.py new file mode 100644 index 00000000..99604fd3 --- /dev/null +++ b/tests/integration/test_topup_untrusted_mint.py @@ -0,0 +1,56 @@ +""" +Integration test for the wallet topup endpoint with a foreign-mint token. + +Tokens are only accepted from trusted mints (primary_mint plus cashu_mints) +and are always redeemed on the mint that issued them. A token from any other +mint is rejected offline, before any network contact with that mint, with a +dedicated error type and code. This replaces the former cross-mint swap +path, so there is no fee-retry behaviour left to exercise here. +""" + +from unittest.mock import AsyncMock, Mock, patch + +import pytest +from httpx import AsyncClient + +from routstr.core.settings import settings + +# Captured at collection time, before the integration_app fixture replaces it +# with the testmint stub (see conftest.py). +from routstr.wallet import recieve_token as _real_recieve_token + +PRIMARY_MINT = "http://localhost:3338" +FOREIGN_MINT = "http://foreign-mint:3338" + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_topup_with_foreign_mint_token_is_rejected_without_mint_contact( + authenticated_client: AsyncClient, +) -> None: + mock_token = Mock() + mock_token.mint = FOREIGN_MINT + mock_token.unit = "sat" + mock_token.amount = 1000 + mock_token.keysets = ["keyset"] + get_wallet = AsyncMock() + + with ( + patch("routstr.wallet.recieve_token", _real_recieve_token), + patch("routstr.wallet.deserialize_token_from_string", return_value=mock_token), + patch("routstr.wallet.get_wallet", get_wallet), + patch.object(settings, "primary_mint", PRIMARY_MINT), + patch.object(settings, "primary_mint_unit", "sat"), + patch.object(settings, "cashu_mints", [PRIMARY_MINT]), + ): + response = await authenticated_client.post( + "/v1/wallet/topup", + params={"cashu_token": "cashuAtest_foreign_token"}, + ) + + assert response.status_code == 400 + error = response.json()["detail"]["error"] + assert error["type"] == "untrusted_mint" + assert error["code"] == "cashu_untrusted_source_mint" + assert FOREIGN_MINT not in error["message"] + get_wallet.assert_not_awaited() diff --git a/tests/integration/test_wallet_authentication.py b/tests/integration/test_wallet_authentication.py index 9dc19947..26407701 100644 --- a/tests/integration/test_wallet_authentication.py +++ b/tests/integration/test_wallet_authentication.py @@ -3,6 +3,7 @@ Integration tests for wallet authentication system including API key generation Tests POST /v1/wallet/topup endpoint and authorization header validation. """ +import asyncio from datetime import datetime, timedelta from typing import Any @@ -113,38 +114,32 @@ async def test_api_key_generation_invalid_token( async def test_duplicate_token_handling( integration_client: AsyncClient, testmint_wallet: Any, db_snapshot: Any ) -> None: - """Test that duplicate tokens return the same API key without double-spending""" - - # Generate a valid token - amount = 500 # 500 sats + amount = 500 token = await testmint_wallet.mint_tokens(amount) - - # First use of token integration_client.headers["Authorization"] = f"Bearer {token}" - response1 = await integration_client.get("/v1/wallet/info") - assert response1.status_code == 200 + + response1, response2 = await asyncio.gather( + integration_client.get("/v1/wallet/info"), + integration_client.get("/v1/wallet/info"), + ) + assert response1.status_code < 500 + assert response2.status_code < 500 + assert response1.status_code == response2.status_code == 200 api_key1 = response1.json()["api_key"] - balance1 = response1.json()["balance"] - - # Capture state after first submission - await db_snapshot.capture() - - # Second use of same token - should return same API key since it's already created - response2 = await integration_client.get("/v1/wallet/info") - assert response2.status_code == 200 api_key2 = response2.json()["api_key"] + balance1 = response1.json()["balance"] balance2 = response2.json()["balance"] - - # Should return the same API key and balance assert api_key1 == api_key2 - assert balance1 == balance2 + assert balance1 == balance2 == amount * 1000 - # Verify no additional database changes + await db_snapshot.capture() + replay = await integration_client.get("/v1/wallet/info") + assert replay.status_code == 200 + assert replay.json()["api_key"] == api_key1 diff = await db_snapshot.diff() assert len(diff["api_keys"]["added"]) == 0 assert len(diff["api_keys"]["modified"]) == 0 - # Original API key should still work with original balance integration_client.headers["Authorization"] = f"Bearer {api_key1}" wallet_response = await integration_client.get("/v1/wallet/") assert wallet_response.status_code == 200 diff --git a/tests/integration/test_wallet_melt_restart.py b/tests/integration/test_wallet_melt_restart.py index 5bb607fa..9d75f542 100644 --- a/tests/integration/test_wallet_melt_restart.py +++ b/tests/integration/test_wallet_melt_restart.py @@ -10,9 +10,12 @@ invalidate them on "paid" or release them on "unpaid". """ from pathlib import Path +from unittest.mock import AsyncMock, patch +import httpx import pytest -from cashu.core.base import Proof +from cashu.core.base import MeltQuote, MeltQuoteState, Proof +from cashu.core.models import PostMeltQuoteResponse from cashu.wallet import crud from cashu.wallet.wallet import Wallet @@ -47,6 +50,20 @@ async def _seed_ambiguous_melt(wallet: Wallet) -> list[Proof]: proofs = [_proof("secret-a"), _proof("secret-b", amount=32)] for proof in proofs: await crud.store_proof(proof, db=wallet.db) + await crud.store_bolt11_melt_quote( + db=wallet.db, + quote=MeltQuote( + quote=QUOTE_ID, + method="bolt11", + request="lnbc1-test", + checking_id="", + unit="sat", + amount=95, + fee_reserve=1, + state=MeltQuoteState.pending, + mint=str(wallet.url), + ), + ) await wallet.set_reserved_for_melt(proofs, reserved=True, quote_id=QUOTE_ID) # cashu's `except` block in melt(): @@ -75,6 +92,34 @@ async def test_melt_recovery_is_findable_by_quote_after_restart( assert all(p.melt_id == QUOTE_ID for p in found) +async def test_paid_reconciliation_invalidates_recovered_proofs_after_restart( + tmp_path: Path, +) -> None: + wallet = await _wallet(tmp_path) + await _seed_ambiguous_melt(wallet) + + restarted = await _wallet(tmp_path) + remote = PostMeltQuoteResponse( + quote=QUOTE_ID, + amount=95, + unit="sat", + request="lnbc1-test", + fee_reserve=1, + state=MeltQuoteState.paid.value, + expiry=None, + payment_preimage="preimage", + ) + with patch( + "cashu.wallet.v1_api.LedgerAPI.get_melt_quote", + new=AsyncMock(return_value=remote), + ): + reconciled = await restarted.get_melt_quote(QUOTE_ID) + + assert reconciled is not None and reconciled.state == MeltQuoteState.paid + assert await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) == [] + assert await crud.get_proofs(db=restarted.db) == [] + + async def test_send_style_reservation_would_not_be_reconcilable( tmp_path: Path, ) -> None: @@ -97,19 +142,64 @@ async def test_send_style_reservation_would_not_be_reconcilable( async def test_unpaid_reconciliation_releases_recovered_proofs_after_restart( tmp_path: Path, ) -> None: - """The full recovery arc: crash, restart, mint says unpaid, funds usable.""" wallet = await _wallet(tmp_path) await _seed_ambiguous_melt(wallet) restarted = await _wallet(tmp_path) - found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) - assert len(found) == 2 + remote = PostMeltQuoteResponse( + quote=QUOTE_ID, + amount=95, + unit="sat", + request="lnbc1-test", + fee_reserve=1, + state=MeltQuoteState.unpaid.value, + expiry=None, + ) + with patch( + "cashu.wallet.v1_api.LedgerAPI.get_melt_quote", + new=AsyncMock(return_value=remote), + ): + reconciled = await restarted.get_melt_quote(QUOTE_ID) - # What get_melt_quote() does on an "unpaid" answer. - await restarted.set_reserved_for_melt(found, reserved=False, quote_id=None) - - released = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) - assert released == [] + assert reconciled is not None and reconciled.state == MeltQuoteState.unpaid + assert await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) == [] all_proofs = await crud.get_proofs(db=restarted.db) assert len(all_proofs) == 2 - assert all(not p.reserved for p in all_proofs) # spendable again + assert all(not p.reserved for p in all_proofs) + + +async def test_pending_and_transport_reconciliation_keep_recovered_reservation( + tmp_path: Path, +) -> None: + wallet = await _wallet(tmp_path) + await _seed_ambiguous_melt(wallet) + restarted = await _wallet(tmp_path) + pending = PostMeltQuoteResponse( + quote=QUOTE_ID, + amount=95, + unit="sat", + request="lnbc1-test", + fee_reserve=1, + state=MeltQuoteState.pending.value, + expiry=None, + ) + + with patch( + "cashu.wallet.v1_api.LedgerAPI.get_melt_quote", + new=AsyncMock(return_value=pending), + ): + reconciled = await restarted.get_melt_quote(QUOTE_ID) + assert reconciled is not None and reconciled.state == MeltQuoteState.pending + found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) + assert len(found) == 2 and all(proof.reserved for proof in found) + + with ( + patch( + "cashu.wallet.v1_api.LedgerAPI.get_melt_quote", + new=AsyncMock(side_effect=httpx.ReadTimeout("mint unavailable")), + ), + pytest.raises(httpx.ReadTimeout), + ): + await restarted.get_melt_quote(QUOTE_ID) + found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) + assert len(found) == 2 and all(proof.reserved for proof in found) diff --git a/tests/integration/test_wallet_refund.py b/tests/integration/test_wallet_refund.py index 1cbd6b4b..1029d073 100644 --- a/tests/integration/test_wallet_refund.py +++ b/tests/integration/test_wallet_refund.py @@ -7,7 +7,7 @@ import asyncio import base64 import json from typing import Any -from unittest.mock import patch +from unittest.mock import AsyncMock, patch import pytest from httpx import AsyncClient @@ -204,7 +204,7 @@ async def test_refund_with_lightning_address( await db_snapshot.capture() # Mock send_to_lnurl function directly - with patch("routstr.balance.send_to_lnurl") as mock_send_to_lnurl: + with patch("routstr.refund.send_to_lnurl") as mock_send_to_lnurl: mock_send_to_lnurl.return_value = { "amount_sent": balance, "unit": "msat", @@ -508,7 +508,7 @@ async def test_mint_unavailability_handling( # Make the send_token method raise a typed mint connection exception. with patch( - "routstr.balance.send_token", + "routstr.refund.send_token", side_effect=MintConnectionError(raw_error), ): response = await authenticated_client.post("/v1/wallet/refund") @@ -622,7 +622,10 @@ async def test_refund_with_expired_key( integration_client.headers["Authorization"] = f"Bearer {api_key}" # Mock the refund to LN address - with patch("routstr.balance.send_to_lnurl") as mock_send_to_lnurl: + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()), + patch("routstr.refund.send_to_lnurl") as mock_send_to_lnurl, + ): mock_send_to_lnurl.return_value = 500 response = await integration_client.post("/v1/wallet/refund") diff --git a/tests/integration/test_wallet_topup.py b/tests/integration/test_wallet_topup.py index 37d59651..3a7bff6f 100644 --- a/tests/integration/test_wallet_topup.py +++ b/tests/integration/test_wallet_topup.py @@ -337,16 +337,13 @@ async def test_topup_during_active_proxy_request( # type: ignore[no-untyped-def @pytest.mark.integration @pytest.mark.asyncio -async def test_maximum_balance_limits( # type: ignore[no-untyped-def] +async def test_large_balance_topup_is_allowed( # type: ignore[no-untyped-def] integration_client: AsyncClient, authenticated_client: AsyncClient, testmint_wallet: Any, integration_session, ) -> None: - """Test if there are any maximum balance limits""" - - # Note: The current implementation doesn't enforce maximum balance limits - # This test verifies large balances are handled correctly + """Test that a large top-up is accepted and reflected in the balance.""" # Get current balance response = await authenticated_client.get("/v1/wallet/") diff --git a/tests/unit/test_admin_logs_by_request_id.py b/tests/unit/test_admin_logs_by_request_id.py new file mode 100644 index 00000000..5391654a --- /dev/null +++ b/tests/unit/test_admin_logs_by_request_id.py @@ -0,0 +1,39 @@ +from typing import Any +from unittest.mock import Mock, patch + +import pytest + +from routstr.core.admin import get_logs_by_request_id_api + + +@pytest.mark.asyncio +async def test_returns_entries_for_request_id_oldest_first() -> None: + entries = [ + {"asctime": "2026-01-01 10:00:02", "message": "done", "request_id": "req-1"}, + {"asctime": "2026-01-01 10:00:01", "message": "start", "request_id": "req-1"}, + ] + + with patch("routstr.core.admin.log_manager") as log_manager: + log_manager.search_logs.return_value = entries + result: dict[str, Any] = await get_logs_by_request_id_api( + request=Mock(), request_id="req-1", date=None, limit=200 + ) + + log_manager.search_logs.assert_called_once_with( + date=None, request_id="req-1", limit=200 + ) + assert result["total"] == 2 + assert result["request_id"] == "req-1" + assert [entry["message"] for entry in result["logs"]] == ["start", "done"] # type: ignore[index,union-attr] + + +@pytest.mark.asyncio +async def test_returns_empty_list_for_unknown_request_id() -> None: + with patch("routstr.core.admin.log_manager") as log_manager: + log_manager.search_logs.return_value = [] + result = await get_logs_by_request_id_api( + request=Mock(), request_id="missing", date="2026-01-01", limit=10 + ) + + assert result["logs"] == [] + assert result["total"] == 0 diff --git a/tests/unit/test_admin_transactions.py b/tests/unit/test_admin_transactions.py index 22508fd6..4ce6ffa0 100644 --- a/tests/unit/test_admin_transactions.py +++ b/tests/unit/test_admin_transactions.py @@ -1,9 +1,14 @@ +from collections.abc import AsyncIterator from contextlib import asynccontextmanager +from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch import pytest +from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession -from routstr.core.admin import get_transactions_api +from routstr.core.admin import _transaction_status, get_transactions_api from routstr.core.db import CashuTransaction @@ -33,3 +38,151 @@ async def test_transactions_api_excludes_internal_sweep_claim_timestamp() -> Non assert response["total"] == 1 assert response["transactions"][0]["token"] == "cashu-token" assert "sweep_started_at" not in response["transactions"][0] + + +@pytest.mark.parametrize( + ("source", "typ", "collected", "swept", "expected"), + [ + ("admin", "out", False, False, "issued"), + ("admin", "out", True, False, "collected"), + ("admin", "out", False, True, "swept"), + # The node redeems an incoming top-up itself, so it is never "issued". + ("admin", "in", False, False, "pending"), + ("admin", "in", True, False, "collected"), + ("x-cashu", "out", False, False, "pending"), + ("x-cashu", "out", True, False, "collected"), + ("apikey", "out", False, True, "swept"), + ], +) +def test_transaction_status( + source: str, typ: str, collected: bool, swept: bool, expected: str +) -> None: + transaction = CashuTransaction( + token="t", + amount=1, + unit="sat", + type=typ, + source=source, + collected=collected, + swept=swept, + ) + assert _transaction_status(transaction) == expected + + +@pytest.fixture +async def session_factory( + tmp_path: Path, +) -> AsyncIterator[async_sessionmaker[AsyncSession]]: + engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'admin.db'}") + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + yield factory + await engine.dispose() + + +ROWS = [ + ("issued-withdrawal", "admin", "out", False, False), + ("collected-withdrawal", "admin", "out", True, False), + ("swept-withdrawal", "admin", "out", False, True), + ("incoming-topup", "admin", "in", False, False), + ("pending-xcashu", "x-cashu", "out", False, False), + ("collected-apikey", "apikey", "out", True, False), + # Nothing stops both terminal flags being set, and the sweeper can reach + # this state when a token it claimed turns out to be already spent. + ("collected-and-swept", "x-cashu", "out", True, True), +] + + +async def _seed(factory: async_sessionmaker[AsyncSession]) -> None: + async with factory() as session: + session.add_all( + [ + CashuTransaction( + id=row_id, + token=row_id, + amount=1, + unit="sat", + source=source, + type=typ, + collected=collected, + swept=swept, + ) + for row_id, source, typ, collected, swept in ROWS + ] + ) + await session.commit() + + +async def _query( + factory: async_sessionmaker[AsyncSession], **kwargs: object +) -> list[dict]: + @asynccontextmanager + async def create_session(): # type: ignore[no-untyped-def] + async with factory() as session: + yield session + + with patch("routstr.core.admin.create_session", create_session): + response = await get_transactions_api(**kwargs) # type: ignore[arg-type] + return response["transactions"] + + +@pytest.mark.asyncio +async def test_response_carries_status_for_every_row( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _seed(session_factory) + + by_id = {tx["id"]: tx["status"] for tx in await _query(session_factory)} + + assert by_id == { + "issued-withdrawal": "issued", + "collected-withdrawal": "collected", + "swept-withdrawal": "swept", + "incoming-topup": "pending", + "pending-xcashu": "pending", + "collected-apikey": "collected", + "collected-and-swept": "swept", + } + + +@pytest.mark.asyncio +async def test_status_filters_partition_rows_by_reported_status( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _seed(session_factory) + unfiltered = {tx["id"]: tx["status"] for tx in await _query(session_factory)} + + filtered: dict[str, str] = {} + for status in ("issued", "collected", "swept", "pending"): + for tx in await _query(session_factory, status=status): + assert tx["status"] == status + filtered[tx["id"]] = status + + assert filtered == unfiltered + + +@pytest.mark.asyncio +async def test_issued_filter_returns_only_outstanding_withdrawals( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _seed(session_factory) + + rows = await _query(session_factory, status="issued") + + assert [tx["id"] for tx in rows] == ["issued-withdrawal"] + + +@pytest.mark.asyncio +async def test_pending_filter_excludes_issued_withdrawals( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _seed(session_factory) + + rows = await _query(session_factory, status="pending") + ids = {tx["id"] for tx in rows} + + # It is uncollected and unswept, so it matched "pending" before it had a + # status of its own. + assert "issued-withdrawal" not in ids + assert ids == {"incoming-topup", "pending-xcashu"} diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index 0dcf751e..f5c875f9 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -499,7 +499,14 @@ def test_create_model_mappings_exact_model_id_beats_forwarded_id_collision( def test_models_endpoint_preserves_catalog_id_when_winner_forwards_elsewhere( monkeypatch: pytest.MonkeyPatch, ) -> None: - """Each catalog row keeps its requested ID while using its routing winner.""" + """Each catalog row keeps its requested ID while using its routing winner. + + ``foo`` is served directly by two providers: the cheaper one prefixes the ID + (``vendor/foo``) and the pricier one exposes the bare ID while forwarding + upstream to ``bar``. Prefix vs bare spelling must not decide the winner, so + the cheaper prefixed provider wins ``foo`` while ``bar`` still appears as its + own catalog row served by the forwarding provider. + """ base_alias = create_test_model( "vendor/foo", prompt_price=0.001, completion_price=0.001 ) @@ -518,9 +525,9 @@ def test_models_endpoint_preserves_catalog_id_when_winner_forwards_elsewhere( disabled_model_keys=set(), ) - assert provider_map["foo"][0] == (redirected_exact, redirect_provider) + assert provider_map["foo"][0] == (base_alias, base_provider) assert unique_models["foo"].id == "foo" - assert unique_models["foo"].upstream_provider_id == "redirect" + assert unique_models["foo"].upstream_provider_id == "base" import routstr.proxy as proxy @@ -876,3 +883,175 @@ def test_create_model_mappings_disables_only_matching_provider() -> None: ) assert [p for _, p in provider_map["same-id"]] == [provider_a] + + +def test_create_model_mappings_prefixed_openrouter_beats_bare_tinfoil_id() -> None: + """Prefix-vs-bare ID spelling must not outrank price for the same model. + + Tinfoil advertises bare model IDs (``gpt-oss-120b``) while OpenRouter keeps + the org prefix (``openai/gpt-oss-120b``). Both serve the same model, so the + cheaper OpenRouter deployment must win the public ``gpt-oss-120b`` catalog + row and route; the bare-ID exact match must not shadow it on spelling alone. + """ + tinfoil_expensive = create_test_model( + "gpt-oss-120b", prompt_price=0.01, completion_price=0.01 + ) + openrouter_cheap = create_test_model( + "openai/gpt-oss-120b", prompt_price=0.001, completion_price=0.001 + ) + tinfoil = create_test_provider( + "tinfoil", + "https://inference.tinfoil.sh/v1", + db_id=1, + models=[tinfoil_expensive], + ) + openrouter = create_test_provider( + "openrouter", + "https://openrouter.ai/api/v1", + db_id=2, + models=[openrouter_cheap], + ) + + # Discovery order should not matter: Tinfoil (non-OpenRouter) is processed + # first, yet the cheaper OpenRouter candidate must still win. + _, provider_map, unique_models = create_model_mappings( + upstreams=[tinfoil, openrouter], + overrides_by_key={}, + disabled_model_keys=set(), + ) + + assert provider_map["gpt-oss-120b"][0] == (openrouter_cheap, openrouter) + assert unique_models["gpt-oss-120b"].upstream_provider_id == "openrouter" + assert unique_models["gpt-oss-120b"].pricing.prompt == 0.001 + + +def test_create_model_mappings_uppercase_prefixed_base_keeps_top_tier() -> None: + """Uppercase prefixed IDs still match the public alias at the direct tier. + + ``Qwen/Qwen2.5-72B`` lowercases to alias ``qwen2.5-72b``; its base name + must be compared case-insensitively so it stays a direct match instead of + falling to the weakest tier and losing to a forwarded alias on spelling. + """ + prefixed_cheap = create_test_model( + "Qwen/Qwen2.5-72B", prompt_price=0.001, completion_price=0.001 + ) + forwarded_expensive = create_test_model( + "deployment-x", prompt_price=0.1, completion_price=0.1 + ) + forwarded_expensive.forwarded_model_id = "qwen2.5-72b" + prefixed_provider = create_test_provider( + "prefixed", "https://prefixed.example/v1", db_id=1, models=[prefixed_cheap] + ) + forwarded_provider = create_test_provider( + "forwarded", "https://forwarded.example/v1", db_id=2, models=[forwarded_expensive] + ) + + _, provider_map, unique_models = create_model_mappings( + upstreams=[forwarded_provider, prefixed_provider], + overrides_by_key={}, + disabled_model_keys=set(), + ) + + assert provider_map["qwen2.5-72b"][0] == (prefixed_cheap, prefixed_provider) + assert unique_models["qwen2.5-72b"].upstream_provider_id == "prefixed" + + +def test_create_model_mappings_excludes_a_malformed_price() -> None: + """A rate that is not a number must not be routable. + + A negative or non-finite rate reads as a real price to every truthiness + check, so the candidate was built into the map and served. The cost + calculation cannot price on such a rate, so every request on the model fell + through to the flat maximum reservation — or, for a negative rate, billed a + negative amount that settlement credits back to the caller. + """ + healthy = create_test_model("healthy-model") + for bad_rate in (float("nan"), float("inf"), -1.0): + broken = create_test_model("broken-model", prompt_price=bad_rate) + provider = create_test_provider( + "custom", + "https://custom.example/v1", + db_id=1, + models=[broken, healthy], + ) + + _, provider_map, unique_models = create_model_mappings( + upstreams=[provider], + overrides_by_key={}, + disabled_model_keys=set(), + ) + + assert "broken-model" not in provider_map, bad_rate + assert "broken-model" not in unique_models, bad_rate + # One unroutable candidate must not cost the provider its other models. + assert "healthy-model" in provider_map, bad_rate + + +def test_create_model_mappings_excludes_an_override_with_a_malformed_price( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """An override row carrying a malformed rate is unroutable too. + + An override replaces the discovered model's price, so a provider whose + catalog is sound still routes at whatever the row says. The guard has to sit + after the override is applied, not before it. + """ + discovered = create_test_model("shared-model") + provider = create_test_provider( + "custom", "https://custom.example/v1", db_id=3, models=[discovered] + ) + override_model = create_test_model("shared-model", prompt_price=float("-inf")) + + monkeypatch.setattr( + "routstr.payment.models._row_to_model", + lambda *args, **kwargs: override_model, + ) + override_row = SimpleNamespace( + id="shared-model", upstream_provider_id=3, enabled=True + ) + + _, provider_map, unique_models = create_model_mappings( + upstreams=[provider], + overrides_by_key={("shared-model", 3): (override_row, 1.0)}, + disabled_model_keys=set(), + ) + + assert "shared-model" not in provider_map + assert "shared-model" not in unique_models + + +def test_create_model_mappings_survives_an_unreadable_override_row( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """One row that cannot be read must not empty the whole routing map. + + Converting an override while walking a provider's catalog let the exception + unwind the entire map build: the node came up routing nothing. The sibling + loop over override-only rows already skips and logs such a row. + """ + broken = create_test_model("broken-model") + healthy = create_test_model("healthy-model") + provider = create_test_provider( + "custom", + "https://custom.example/v1", + db_id=5, + models=[broken, healthy], + ) + + def raising_row_to_model(row: Any, *args: Any, **kwargs: Any) -> Model: + raise ValueError("value is not a valid float") + + monkeypatch.setattr("routstr.payment.models._row_to_model", raising_row_to_model) + override_row = SimpleNamespace( + id="broken-model", upstream_provider_id=5, enabled=True + ) + + _, provider_map, unique_models = create_model_mappings( + upstreams=[provider], + overrides_by_key={("broken-model", 5): (override_row, 1.0)}, + disabled_model_keys=set(), + ) + + assert "broken-model" not in provider_map + assert "healthy-model" in provider_map + assert "healthy-model" in unique_models diff --git a/tests/unit/test_auth_cashu.py b/tests/unit/test_auth_cashu.py index 29421231..b221fbef 100644 --- a/tests/unit/test_auth_cashu.py +++ b/tests/unit/test_auth_cashu.py @@ -88,7 +88,7 @@ async def test_failed_first_cashu_redemption_rolls_back_empty_api_key( httpx.ConnectError("All connection attempts failed"), 503, "mint_unreachable", - "Cashu mint is unreachable", + "Cashu mint is unreachable; retry later", "cashu_mint_unreachable", ), ( @@ -96,7 +96,7 @@ async def test_failed_first_cashu_redemption_rolls_back_empty_api_key( MintConnectionError("connect to http://mint:3338 refused"), 503, "mint_unreachable", - "Cashu mint is unreachable", + "Cashu mint is unreachable; retry later", "cashu_mint_unreachable", ), ( @@ -104,16 +104,16 @@ async def test_failed_first_cashu_redemption_rolls_back_empty_api_key( _value_error_wrapping_transport(), 503, "mint_unreachable", - "Cashu mint is unreachable", + "Cashu mint is unreachable; retry later", "cashu_mint_unreachable", ), ( # asyncio.TimeoutError is builtin TimeoutError on 3.11+. TimeoutError("Timed out connecting to Cashu mint http://mint:3338"), 503, - "mint_unreachable", - "Cashu mint is unreachable", - "cashu_mint_unreachable", + "mint_timeout", + "Cashu mint did not respond in time; retry later", + "cashu_mint_timeout", ), ( ValueError( @@ -134,15 +134,6 @@ async def test_failed_first_cashu_redemption_rolls_back_empty_api_key( "Token value is too small to cover swap fees", "cashu_token_swap_fees_exceed_amount", ), - ( - ValueError( - "Failed to melt token from foreign mint http://foreign:3338: boom" - ), - 422, - "mint_error", - "Failed to swap token from foreign mint", - "cashu_foreign_mint_swap_failed", - ), ( ValueError("could not decode token"), 400, diff --git a/tests/unit/test_auto_topup.py b/tests/unit/test_auto_topup.py index 0db7408f..349dc9ef 100644 --- a/tests/unit/test_auto_topup.py +++ b/tests/unit/test_auto_topup.py @@ -1,14 +1,16 @@ import json +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest -from routstr.core.db import CashuTransaction from routstr.upstream.auto_topup import ( _check_and_topup, _parse_ppq_request_id, _run_auto_topup_cycle, validate_ppq_auto_topup_settings, + validate_routstr_auto_topup_settings, ) from routstr.upstream.ppqai import PPQAIUpstreamProvider from routstr.wallet import Bolt11PaymentAmbiguous, Bolt11PaymentNotAttempted @@ -46,35 +48,19 @@ def _row() -> MagicMock: return row -class _Session: - def __init__(self, transaction: CashuTransaction) -> None: - self.transaction = transaction - self.commit = AsyncMock() - - async def __aenter__(self) -> "_Session": - return self - - async def __aexit__(self, *args: object) -> None: - return None - - async def exec(self, query: object) -> MagicMock: - result = MagicMock() - result.first.return_value = self.transaction - return result - - def add(self, transaction: CashuTransaction) -> None: - self.transaction = transaction - - @pytest.mark.asyncio -async def test_auto_topup_persists_before_sending_and_marks_success_collected() -> None: +async def test_auto_topup_refuses_invalid_settings_before_touching_the_wallet() -> None: provider = MagicMock() - provider.get_balance = AsyncMock(return_value=0) - provider.topup = AsyncMock(return_value={"balance": 50}) - transaction = CashuTransaction( - token="cashu-token", amount=50, unit="sat", source="auto_topup" + provider.get_balance = AsyncMock() + row = _row() + row.provider_settings = json.dumps( + { + "auto_topup": True, + "topup_threshold": 100, + "topup_amount_limit": 10**9, + "topup_mint_url": "https://mint.test", + } ) - session = _Session(transaction) with ( patch( @@ -82,99 +68,71 @@ async def test_auto_topup_persists_before_sending_and_marks_success_collected() return_value=provider, ), patch( - "routstr.upstream.auto_topup.send_token", - AsyncMock(return_value="cashu-token"), + "routstr.upstream.auto_topup.send_token_from_owner_locked", AsyncMock() + ) as send, + ): + await _check_and_topup(row) + + provider.get_balance.assert_not_awaited() + send.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_routstr_outgoing_audit_is_persisted_under_wallet_guard() -> None: + provider = MagicMock() + provider.get_balance = AsyncMock(return_value=0.0) + provider.topup = AsyncMock(return_value={"error": "stop after send"}) + inside_guard = False + + @asynccontextmanager + async def guard() -> AsyncIterator[None]: + nonlocal inside_guard + inside_guard = True + try: + yield + finally: + inside_guard = False + + async def send(*_args: object) -> str: + assert inside_guard + return "cashu-token" + + async def persist(*_args: object, **_kwargs: object) -> None: + assert inside_guard + + with ( + patch( + "routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row", + return_value=provider, ), patch( - "routstr.upstream.auto_topup.store_cashu_transaction", - AsyncMock(return_value=True), - ) as store, + "routstr.upstream.auto_topup._reconcile_routstr_state", + AsyncMock(return_value=False), + ), + patch( + "routstr.upstream.auto_topup._routstr_spent_last_24h_sats", + AsyncMock(return_value=0), + ), + patch( + "routstr.upstream.auto_topup._claim_routstr_topup", + AsyncMock(return_value="operation-1"), + ), + patch("routstr.upstream.auto_topup.wallet_operation_guard", side_effect=guard), + patch( + "routstr.upstream.auto_topup.send_token_from_owner_locked", + side_effect=send, + ), + patch( + "routstr.upstream.auto_topup._persist_routstr_token_and_mark_sent", + side_effect=persist, + ), patch( "routstr.upstream.auto_topup.token_mint_url", - return_value="https://fallback-mint.test", + return_value="https://mint.test", ), - patch("routstr.upstream.auto_topup.create_session", return_value=session), ): await _check_and_topup(_row()) - store.assert_awaited_once_with( - token="cashu-token", - amount=50, - unit="sat", - mint_url="https://fallback-mint.test", - typ="out", - collected=False, - source="auto_topup", - ) - provider.topup.assert_awaited_once_with("cashu-token") - assert transaction.collected is True - session.commit.assert_awaited_once() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("outcome", [{"error": "rejected"}, RuntimeError("network")]) -async def test_auto_topup_failure_leaves_persisted_token_uncollected( - outcome: object, -) -> None: - provider = MagicMock() - provider.get_balance = AsyncMock(return_value=0) - provider.topup = AsyncMock( - side_effect=outcome if isinstance(outcome, Exception) else None, - return_value=outcome, - ) - - with ( - patch( - "routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row", - return_value=provider, - ), - patch( - "routstr.upstream.auto_topup.send_token", - AsyncMock(return_value="cashu-token"), - ), - patch( - "routstr.upstream.auto_topup.store_cashu_transaction", - AsyncMock(return_value=True), - ), - patch("routstr.upstream.auto_topup.create_session") as create_session, - ): - if isinstance(outcome, Exception): - with pytest.raises(RuntimeError): - await _check_and_topup(_row()) - else: - await _check_and_topup(_row()) - - create_session.assert_not_called() - - -@pytest.mark.asyncio -async def test_auto_topup_does_not_send_untracked_token() -> None: - provider = MagicMock() - provider.get_balance = AsyncMock(return_value=0) - provider.topup = AsyncMock() - with ( - patch( - "routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row", - return_value=provider, - ), - patch( - "routstr.upstream.auto_topup.send_token", - AsyncMock(return_value="cashu-token"), - ), - patch( - "routstr.upstream.auto_topup.store_cashu_transaction", - AsyncMock(side_effect=RuntimeError("database unavailable")), - ), - patch( - "routstr.upstream.auto_topup.release_token_reservation", - AsyncMock(), - ) as reclaim, - ): - await _check_and_topup(_row()) - - reclaim.assert_awaited_once_with("cashu-token") - provider.topup.assert_not_awaited() - def _ppq_row() -> MagicMock: row = MagicMock() @@ -546,6 +504,30 @@ async def test_ppq_auto_topup_skips_when_balance_meets_threshold() -> None: provider.initiate_topup.assert_not_awaited() +@pytest.mark.asyncio +async def test_ppq_auto_topup_requires_two_below_threshold_reads() -> None: + provider = MagicMock() + provider.get_balance = AsyncMock(side_effect=[2.5, 5.0]) + provider.initiate_topup = AsyncMock() + + with ( + patch( + "routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row", + return_value=provider, + ), + patch( + "routstr.upstream.auto_topup._reconcile_ppq_state", + AsyncMock(return_value=False), + ), + patch("routstr.upstream.auto_topup._claim_ppq_topup", AsyncMock()) as claim, + ): + await _check_and_topup(_ppq_row()) + + assert provider.get_balance.await_count == 2 + claim.assert_not_awaited() + provider.initiate_topup.assert_not_awaited() + + @pytest.mark.asyncio async def test_ppq_auto_topup_skips_when_daily_spend_cap_reached() -> None: provider = MagicMock() @@ -734,3 +716,47 @@ def test_ppq_auto_topup_settings_validation_survives_huge_json_integers() -> Non {"auto_topup": True, "topup_threshold": 10**400, "topup_amount_limit": 10} ) assert problem is not None and "threshold" in problem + + +def _routstr_settings(**overrides: object) -> dict: + settings = { + "auto_topup": True, + "topup_threshold": 1, + "topup_amount_limit": 50, + "topup_mint_url": "https://mint.test", + } + settings.update(overrides) + return settings + + +@pytest.mark.parametrize( + ("settings", "expected"), + [ + ({"auto_topup": False, "topup_threshold": -1}, None), + (_routstr_settings(), None), + (_routstr_settings(topup_threshold=None), "threshold"), + (_routstr_settings(topup_threshold=True), "threshold"), + (_routstr_settings(topup_threshold=float("inf")), "threshold"), + (_routstr_settings(topup_amount_limit=0), "positive"), + (_routstr_settings(topup_amount_limit=True), "positive"), + (_routstr_settings(topup_amount_limit=1.5), "whole number"), + (_routstr_settings(topup_amount_limit=10**9), "between"), + (_routstr_settings(topup_mint_url=""), "mint URL"), + (_routstr_settings(topup_mint_url=True), "mint URL"), + ], +) +def test_routstr_auto_topup_settings_validation( + settings: dict, expected: str | None +) -> None: + problem = validate_routstr_auto_topup_settings(settings) + if expected is None: + assert problem is None + else: + assert problem is not None and expected in problem + + +def test_routstr_auto_topup_settings_validation_survives_huge_json_integers() -> None: + problem = validate_routstr_auto_topup_settings( + _routstr_settings(topup_amount_limit=10**400) + ) + assert problem is not None and "positive" in problem diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index 93ac02b8..cb262d70 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -20,7 +20,9 @@ def _make_cashu_tx( swept: bool = False, collected: bool = False, ) -> CashuTransaction: - tx = CashuTransaction(token=token, amount=amount, unit=unit, type=type, request_id=request_id) + tx = CashuTransaction( + token=token, amount=amount, unit=unit, type=type, request_id=request_id + ) tx.swept = swept tx.collected = collected return tx @@ -35,19 +37,32 @@ def _exec_result(tx: CashuTransaction | None) -> MagicMock: def _update_result(rowcount: int) -> MagicMock: result = MagicMock() result.rowcount = rowcount + # Claim and ledger lookups share this stubbed session; an empty row set + # means the key has no prior refund to replay, report, or order after. + result.first.return_value = None + result.one.return_value = None return result @pytest.mark.asyncio async def test_refund_x_cashu_returns_token() -> None: x_cashu_token = "cashuAtest_token_value" - in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-abc") - out_tx = _make_cashu_tx(token="cashuArefund_token", amount=1000, unit="msat", type="out", request_id="req-abc") + in_tx = _make_cashu_tx( + token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-abc" + ) + out_tx = _make_cashu_tx( + token="cashuArefund_token", + amount=1000, + unit="msat", + type="out", + request_id="req-abc", + ) session = MagicMock() session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() result = await refund_wallet_endpoint( authorization="Bearer sk-somekey", @@ -66,13 +81,22 @@ async def test_refund_x_cashu_returns_token() -> None: @pytest.mark.asyncio async def test_refund_x_cashu_sat_unit() -> None: x_cashu_token = "cashuAsat_token" - in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="sat", type="in", request_id="req-sat") - out_tx = _make_cashu_tx(token="cashuArefund_sat", amount=500, unit="sat", type="out", request_id="req-sat") + in_tx = _make_cashu_tx( + token=x_cashu_token, amount=0, unit="sat", type="in", request_id="req-sat" + ) + out_tx = _make_cashu_tx( + token="cashuArefund_sat", + amount=500, + unit="sat", + type="out", + request_id="req-sat", + ) session = MagicMock() session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() result = await refund_wallet_endpoint( authorization="Bearer sk-somekey", @@ -124,6 +148,7 @@ async def test_refund_x_cashu_pending_raises_425() -> None: session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(None)]) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -167,8 +192,21 @@ async def test_refund_x_cashu_in_tx_without_request_id_raises_404() -> None: async def test_refund_x_cashu_swept_raises_410() -> None: from fastapi import HTTPException - in_tx = _make_cashu_tx(token="cashuAswept_token", amount=0, unit="msat", type="in", request_id="req-swept") - out_tx = _make_cashu_tx(token="cashuAswept", amount=100, unit="msat", type="out", request_id="req-swept", swept=True) + in_tx = _make_cashu_tx( + token="cashuAswept_token", + amount=0, + unit="msat", + type="in", + request_id="req-swept", + ) + out_tx = _make_cashu_tx( + token="cashuAswept", + amount=100, + unit="msat", + type="out", + request_id="req-swept", + swept=True, + ) session = MagicMock() session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) @@ -208,7 +246,6 @@ def _make_api_key( refund_currency: str | None = "sat", refund_mint_url: str | None = "https://mint.example.com", refund_address: str | None = None, - parent_key_hash: str | None = None, ) -> ApiKey: key = ApiKey(hashed_key="testhash") key.balance = balance @@ -216,7 +253,6 @@ def _make_api_key( key.refund_currency = refund_currency key.refund_mint_url = refund_mint_url key.refund_address = refund_address - key.parent_key_hash = parent_key_hash key.total_spent = 0 key.total_requests = 0 return key @@ -241,10 +277,12 @@ async def test_apikey_refund_returns_persisted_token_after_cache_loss() -> None: session.exec = AsyncMock(return_value=_exec_result(refund_tx)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance.send_token", AsyncMock()) as mock_send_token, + patch("routstr.refund.latest_open", AsyncMock(return_value=None)), + patch("routstr.refund.latest_terminal", AsyncMock(return_value=None)), + patch("routstr.refund.send_token", AsyncMock()) as mock_send_token, ): result = await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -279,8 +317,12 @@ async def test_apikey_refund_rejects_persisted_token_after_sweep() -> None: session.exec = AsyncMock(return_value=_exec_result(refund_tx)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() - with patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)): + with ( + patch("routstr.refund.latest_open", AsyncMock(return_value=None)), + patch("routstr.refund.latest_terminal", AsyncMock(return_value=None)), + ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -304,13 +346,14 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), - patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), - patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store, - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), + patch("routstr.refund.renew_lease", AsyncMock()), + patch( + "routstr.refund.send_token", AsyncMock(return_value=refund_token) + ) as mock_send_token, + patch("routstr.refund.store_cashu_transaction", AsyncMock()) as mock_store, ): result = await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -318,6 +361,7 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No session=session, ) + mock_send_token.assert_awaited_once() assert isinstance(result, dict) assert result["token"] == refund_token @@ -339,14 +383,15 @@ async def test_apikey_refund_logs_token() -> None: session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), - patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), - patch("routstr.balance.logger") as mock_logger, + patch("routstr.refund.renew_lease", AsyncMock()), + patch( + "routstr.refund.send_token", AsyncMock(return_value=refund_token) + ) as mock_send_token, + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + patch("routstr.refund.logger") as mock_logger, ): await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -354,12 +399,13 @@ async def test_apikey_refund_logs_token() -> None: session=session, ) + mock_send_token.assert_awaited_once() calls = [str(c) for c in mock_logger.info.call_args_list] - assert any("cashu token issued" in c for c in calls) + assert any("refund paid" in c for c in calls) @pytest.mark.asyncio -async def test_apikey_refund_log_includes_path() -> None: +async def test_apikey_refund_log_identifies_the_claim() -> None: key = _make_api_key(balance=5000, refund_currency="sat") refund_token = "cashuApath_token" @@ -368,14 +414,15 @@ async def test_apikey_refund_log_includes_path() -> None: session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), - patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), - patch("routstr.balance.logger") as mock_logger, + patch("routstr.refund.renew_lease", AsyncMock()), + patch( + "routstr.refund.send_token", AsyncMock(return_value=refund_token) + ) as mock_send_token, + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + patch("routstr.refund.logger") as mock_logger, ): await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -383,14 +430,16 @@ async def test_apikey_refund_log_includes_path() -> None: session=session, ) - # Find the "cashu token issued" call and verify extra contains the path - token_issued_calls = [ - c for c in mock_logger.info.call_args_list - if c.args and "cashu token issued" in c.args[0] + mock_send_token.assert_awaited_once() + paid_calls = [ + c + for c in mock_logger.info.call_args_list + if c.args and "refund paid" in c.args[0] ] - assert len(token_issued_calls) == 1 - extra = token_issued_calls[0].kwargs.get("extra", {}) - assert extra.get("path") == "/v1/wallet/refund" + assert len(paid_calls) == 1 + extra = paid_calls[0].kwargs.get("extra", {}) + assert extra.get("method") == "cashu" + assert extra.get("refund_id") @pytest.mark.asyncio @@ -405,15 +454,13 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None: # Debit returns rowcount=0 → balance changed concurrently session.exec = AsyncMock(return_value=_update_result(0)) session.commit = AsyncMock() + session.rollback = AsyncMock() mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted") with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), - patch("routstr.balance.send_token", mock_send_token), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), + patch("routstr.refund.send_token", mock_send_token), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -433,6 +480,7 @@ async def test_credit_balance_stores_apikey_transaction_history() -> None: session = MagicMock() session.exec = AsyncMock(return_value=_update_result(1)) session.commit = AsyncMock() + session.rollback = AsyncMock() session.refresh = AsyncMock() with ( @@ -466,19 +514,19 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None: # First exec call = debit (succeeds), second = restore session = MagicMock() session.get = AsyncMock(return_value=key) - session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)]) + # claim lookup, claim ordering, debit, claim close, balance restore + session.exec = AsyncMock(side_effect=[_update_result(1)] * 5) session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), + patch("routstr.refund.renew_lease", AsyncMock()), patch( - "routstr.balance.send_token", + "routstr.refund.send_token", AsyncMock(side_effect=MintConnectionError("raw mint outage detail")), - ), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), - patch("routstr.balance.logger"), + ) as mock_send_token, + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + patch("routstr.refund.logger"), ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -487,11 +535,12 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None: session=session, ) + mock_send_token.assert_awaited_once() assert exc_info.value.status_code == 503 assert exc_info.value.detail == "Mint service unavailable" assert "raw mint outage detail" not in exc_info.value.detail - # Verify two exec calls: debit + restore - assert session.exec.await_count == 2 + # claim lookup, claim ordering, debit, claim close, balance restore + assert session.exec.await_count == 5 @pytest.mark.asyncio @@ -504,16 +553,18 @@ async def test_apikey_refund_generic_failure_is_sanitized_500() -> None: session = MagicMock() session.get = AsyncMock(return_value=key) - session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)]) + # claim lookup, claim ordering, debit, claim close, balance restore + session.exec = AsyncMock(side_effect=[_update_result(1)] * 5) session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), - patch("routstr.balance.send_token", AsyncMock(side_effect=RuntimeError(raw_error))), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), - patch("routstr.balance.logger"), + patch("routstr.refund.renew_lease", AsyncMock()), + patch( + "routstr.refund.send_token", AsyncMock(side_effect=RuntimeError(raw_error)) + ) as mock_send_token, + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + patch("routstr.refund.logger"), ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -522,10 +573,11 @@ async def test_apikey_refund_generic_failure_is_sanitized_500() -> None: session=session, ) + mock_send_token.assert_awaited_once() assert exc_info.value.status_code == 500 assert exc_info.value.detail == "Refund failed" assert raw_error not in exc_info.value.detail - assert session.exec.await_count == 2 + assert session.exec.await_count == 5 # --------------------------------------------------------------------------- @@ -583,6 +635,7 @@ async def test_refund_unknown_sk_bearer_returns_401() -> None: # --- Topup redemption error taxonomy (POST /v1/wallet/topup) ------------------ + def _envelope(exc: HTTPException) -> dict: """Extract the error object from a top-up HTTPException.""" detail = exc.detail @@ -593,14 +646,31 @@ def _envelope(exc: HTTPException) -> dict: @pytest.mark.asyncio @pytest.mark.parametrize( - "error", + ("error", "expected_type", "expected_code", "expected_message"), [ - httpx.ConnectError("All connection attempts failed"), - MintConnectionError("connect to mint refused"), - TimeoutError("timed out connecting to mint"), + ( + httpx.ConnectError("All connection attempts failed"), + "mint_unreachable", + "cashu_mint_unreachable", + "Cashu mint is unreachable; retry later", + ), + ( + MintConnectionError("connect to mint refused"), + "mint_unreachable", + "cashu_mint_unreachable", + "Cashu mint is unreachable; retry later", + ), + ( + TimeoutError("timed out connecting to mint"), + "mint_timeout", + "cashu_mint_timeout", + "Cashu mint did not respond in time; retry later", + ), ], ) -async def test_topup_mint_unreachable_returns_503(error: Exception) -> None: +async def test_topup_mint_unreachable_returns_503( + error: Exception, expected_type: str, expected_code: str, expected_message: str +) -> None: """A down mint must surface 503 (retryable), not 400 or 500 — the token is fine, so the client should retry once the mint recovers.""" from fastapi import HTTPException @@ -609,7 +679,6 @@ async def test_topup_mint_unreachable_returns_503(error: Exception) -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)), ): with pytest.raises(HTTPException) as exc_info: @@ -619,13 +688,15 @@ async def test_topup_mint_unreachable_returns_503(error: Exception) -> None: assert exc_info.value.status_code == 503 err = _envelope(exc_info.value) - assert err["type"] == "mint_unreachable" - assert err["code"] == "cashu_mint_unreachable" - assert err["message"] == "Cashu mint is unreachable" + assert err["type"] == expected_type + assert err["code"] == expected_code + assert err["message"] == expected_message @pytest.mark.asyncio -async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible() -> None: +async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible() -> ( + None +): from fastapi import HTTPException from routstr.wallet import SourceMintConnectionError @@ -635,7 +706,6 @@ async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible error = SourceMintConnectionError("Issuing Cashu mint is unreachable") with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)), ): with pytest.raises(HTTPException) as exc_info: @@ -647,7 +717,7 @@ async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible err = _envelope(exc_info.value) assert err["type"] == "mint_unreachable" assert err["code"] == "cashu_source_mint_unreachable" - assert "cannot be redeemed at another mint" in err["message"] + assert "retry later" in err["message"] @pytest.mark.asyncio @@ -660,7 +730,6 @@ async def test_topup_already_spent_still_returns_400() -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.credit_balance", AsyncMock(side_effect=ValueError("Token already spent")), @@ -688,11 +757,12 @@ async def test_topup_zero_value_returns_400_zero_value_message() -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.credit_balance", AsyncMock( - side_effect=ValueError("Redeemed token amount must be positive, got 0 msats") + side_effect=ValueError( + "Redeemed token amount must be positive, got 0 msats" + ) ), ), ): @@ -720,7 +790,6 @@ async def test_topup_token_consumed_returns_500() -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.credit_balance", AsyncMock(side_effect=TokenConsumedError("credit failed")), @@ -754,21 +823,12 @@ async def test_topup_token_consumed_returns_500() -> None: "Token value is too small to cover swap fees", ), ( - ValueError( - "Token amount (5 sat) is insufficient to cover melt fees." - ), + ValueError("Token amount (5 sat) is insufficient to cover melt fees."), 422, "mint_error", "cashu_token_swap_fees_exceed_amount", "Token value is too small to cover swap fees", ), - ( - ValueError("Failed to melt token from foreign mint http://m: boom"), - 422, - "mint_error", - "cashu_foreign_mint_swap_failed", - "Failed to swap token from foreign mint", - ), ], ) async def test_topup_fee_and_swap_failures_return_422( @@ -786,7 +846,6 @@ async def test_topup_fee_and_swap_failures_return_422( session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)), ): with pytest.raises(HTTPException) as exc_info: @@ -811,7 +870,6 @@ async def test_topup_unexpected_non_valueerror_returns_500() -> None: session = MagicMock() with ( - patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), patch( "routstr.balance.credit_balance", AsyncMock(side_effect=RuntimeError("db exploded")), @@ -840,17 +898,17 @@ async def test_apikey_refund_ambiguous_melt_does_not_restore_balance() -> None: session = MagicMock() session.get = AsyncMock(return_value=key) - session.exec = AsyncMock(return_value=MagicMock(rowcount=1)) + session.exec = AsyncMock(return_value=_update_result(1)) session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), patch( - "routstr.balance.send_to_lnurl", + "routstr.refund.send_to_lnurl", AsyncMock(side_effect=MeltOutcomeAmbiguousError("outcome is ambiguous")), ), - patch("routstr.balance._restore_balance", AsyncMock()) as mock_restore, + patch("routstr.refund.release", AsyncMock()) as mock_restore, + patch("routstr.refund.get_lnurl_data", AsyncMock()), ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -872,17 +930,17 @@ async def test_apikey_refund_clean_failure_still_restores_balance() -> None: session = MagicMock() session.get = AsyncMock(return_value=key) - session.exec = AsyncMock(return_value=MagicMock(rowcount=1)) + session.exec = AsyncMock(return_value=_update_result(1)) session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), patch( - "routstr.balance.send_to_lnurl", + "routstr.refund.send_to_lnurl", AsyncMock(side_effect=RuntimeError("mint rejected melt")), ), - patch("routstr.balance._restore_balance", AsyncMock()) as mock_restore, + patch("routstr.refund.release", AsyncMock()) as mock_restore, + patch("routstr.refund.get_lnurl_data", AsyncMock()), ): with pytest.raises(HTTPException): await refund_wallet_endpoint( diff --git a/tests/unit/test_cashu_httpx_compat.py b/tests/unit/test_cashu_httpx_compat.py new file mode 100644 index 00000000..3e0162d7 --- /dev/null +++ b/tests/unit/test_cashu_httpx_compat.py @@ -0,0 +1,126 @@ +"""cashu 0.20.x builds its mint client with the `proxies` kwarg httpx removed in +0.28. These tests pin the shim that keeps every wallet call working.""" + +import asyncio + +import httpx +import pytest +from cashu.wallet import v1_api +from httpx import AsyncClient + +from routstr.cashu_compat import ( + _ProxiesCompatAsyncClient, + _single_proxy, + install_cashu_httpx_shim, +) + + +def test_httpx_no_longer_accepts_proxies() -> None: + """The premise of the shim: plain httpx rejects what cashu passes.""" + with pytest.raises(TypeError): + httpx.AsyncClient(proxies={}) # type: ignore[call-arg] + + +@pytest.mark.parametrize( + "proxies, expected", + [ + ({}, None), + (None, None), + ({"all://": "socks5://localhost:9050"}, "socks5://localhost:9050"), + ("socks5://localhost:9050", "socks5://localhost:9050"), + ({"http://": "http://p:1", "https://": "http://p:1"}, "http://p:1"), + ], +) +def test_single_proxy_collapses_cashu_mappings( + proxies: object, expected: str | None +) -> None: + assert _single_proxy(proxies) == expected + + +def test_single_proxy_fails_closed_on_unrepresentable_mapping() -> None: + """Never silently drop a proxy: that would send mint traffic direct.""" + with pytest.raises(ValueError, match="cannot represent proxies"): + _single_proxy({"http://": "http://a:1", "https://": "http://b:2"}) + + +@pytest.mark.parametrize("proxies", [{}, {"all://": "socks5://localhost:9050"}]) +async def test_compat_client_accepts_proxies(proxies: dict) -> None: + async with _ProxiesCompatAsyncClient( + proxies=proxies, base_url="http://mint.test" + ) as client: + assert isinstance(client, httpx.AsyncClient) + + +async def test_empty_proxy_map_disables_environment_proxies( + monkeypatch: pytest.MonkeyPatch, +) -> None: + proxy_url = "http://127.0.0.1:1" + for name in ( + "HTTP_PROXY", + "HTTPS_PROXY", + "ALL_PROXY", + "http_proxy", + "https_proxy", + "all_proxy", + ): + monkeypatch.setenv(name, proxy_url) + monkeypatch.setenv("NO_PROXY", "") + monkeypatch.setenv("no_proxy", "") + + async def respond(_: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + await asyncio.sleep(0) + writer.write(b"HTTP/1.1 204 No Content\r\nContent-Length: 0\r\n\r\n") + await writer.drain() + writer.close() + await writer.wait_closed() + + server = await asyncio.start_server(respond, "127.0.0.1", 0) + try: + port = server.sockets[0].getsockname()[1] + url = f"http://127.0.0.1:{port}/" + + async with _ProxiesCompatAsyncClient(proxies={}, timeout=1) as client: + assert (await client.get(url)).status_code == 204 + + async with _ProxiesCompatAsyncClient( + proxies={}, trust_env=True, timeout=1 + ) as client: + with pytest.raises(httpx.ConnectError): + await client.get(url) + finally: + server.close() + await server.wait_closed() + + +async def test_cashu_decorator_builds_a_client_after_shim() -> None: + """The real cashu decorator — the code path every mint call goes through.""" + install_cashu_httpx_shim() + + class _Ledger: + url = "http://mint.test/" + # cashu's decorator assigns the client here; alias avoids shadowing. + httpx: AsyncClient + + @v1_api.async_set_httpx_client # type: ignore[misc] + async def call(self) -> AsyncClient: + return self.httpx + + client: httpx.AsyncClient = await _Ledger().call() + try: + assert isinstance(client, httpx.AsyncClient) + assert str(client.base_url) == "http://mint.test" + finally: + await client.aclose() + + +def test_install_is_idempotent() -> None: + install_cashu_httpx_shim() + patched = v1_api.httpx + install_cashu_httpx_shim() + assert v1_api.httpx is patched + + +def test_shim_namespace_passes_through_other_httpx_attributes() -> None: + install_cashu_httpx_shim() + assert v1_api.httpx.Response is httpx.Response + assert v1_api.httpx.AsyncClient is _ProxiesCompatAsyncClient diff --git a/tests/unit/test_cashu_untrusted_source_mint.py b/tests/unit/test_cashu_untrusted_source_mint.py new file mode 100644 index 00000000..0129d25f --- /dev/null +++ b/tests/unit/test_cashu_untrusted_source_mint.py @@ -0,0 +1,345 @@ +"""Tokens issued by an untrusted mint are rejected before any mint contact. + +A client-supplied Cashu token names its own mint. Every redemption path used +to load that mint's keysets under ``wallet_operation_guard`` with the full +timeout-retry window, so a silent mint could hold the shared wallet lock for +minutes per request from unauthenticated endpoints. Now the mint must be +``primary_mint`` or one of ``cashu_mints``; anything else fails offline with a +dedicated error type and code. +""" + +from contextlib import ExitStack, contextmanager +from types import SimpleNamespace +from typing import AsyncGenerator, Iterator, cast +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +from fastapi import HTTPException +from sqlalchemy.ext.asyncio import create_async_engine +from sqlalchemy.pool import StaticPool +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.auth import validate_bearer_key +from routstr.core.settings import settings +from routstr.mint import MintCooldownError +from routstr.payment.helpers import check_token_balance +from routstr.wallet import ( + SourceMintConnectionError, + TokenConsumedError, + UntrustedSourceMintError, + classify_redemption_error, + is_mint_timeout, + is_trusted_source_mint, + recieve_token, + resolve_trusted_source_mint, +) + +PRIMARY = "http://primary:3338" +SECONDARY = "http://secondary:3338" +UNTRUSTED = "http://evil:3338" + + +@pytest.fixture +async def session() -> AsyncGenerator[AsyncSession, None]: + engine = create_async_engine( + "sqlite+aiosqlite://", + poolclass=StaticPool, + connect_args={"check_same_thread": False}, + ) + 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() + + +@contextmanager +def _trusted_mints() -> Iterator[None]: + with ExitStack() as stack: + stack.enter_context(patch.object(settings, "primary_mint", PRIMARY)) + stack.enter_context(patch.object(settings, "cashu_mints", [SECONDARY])) + yield + + +def _token(mint: str) -> SimpleNamespace: + return SimpleNamespace(mint=mint, unit="sat", amount=100, keysets=["k"]) + + +def test_is_trusted_source_mint() -> None: + with _trusted_mints(): + assert is_trusted_source_mint(PRIMARY) + assert is_trusted_source_mint(SECONDARY) + assert not is_trusted_source_mint(UNTRUSTED) + + +@pytest.mark.parametrize( + "configured,token_mint", + [ + ("https://mint.example", "https://mint.example/"), + ("https://mint.example/", "https://mint.example"), + ("https://mint.example", "https://mint.example///"), + ("https://mint.example", "HTTPS://MINT.EXAMPLE"), + ("HTTPS://Mint.Example", "https://mint.example"), + ("https://mint.example", "https://mint.example:443"), + ("https://mint.example:443", "https://mint.example"), + ("http://mint.example", "http://mint.example:80"), + (" https://mint.example/ ", "https://mint.example"), + ("https://mint.example/Bitcoin", "https://mint.example/Bitcoin/"), + ], +) +def test_trusted_mint_matching_ignores_cosmetic_url_differences( + configured: str, token_mint: str +) -> None: + with patch.object(settings, "primary_mint", configured): + with patch.object(settings, "cashu_mints", []): + assert is_trusted_source_mint(token_mint) + + +@pytest.mark.parametrize( + "token_mint", + [ + "https://mint.example@evil.example", + "https://mint.example:pw@evil.example", + "https://mint.example.evil.example", + "https://evil.example/?x=https://mint.example", + "https://evil.example#https://mint.example", + "https://mint.example:8443", + "http://mint.example", + "https://mint.example.", + "https://mint.example/bitcoin", + "https://evil.example", + ], +) +def test_trusted_mint_matching_never_folds_onto_another_host(token_mint: str) -> None: + """Normalization must not become an accept-bypass: only cosmetic spelling + differences may fold, never anything that can resolve somewhere else.""" + with patch.object(settings, "primary_mint", "https://mint.example/Bitcoin"): + with patch.object(settings, "cashu_mints", ["https://mint.example"]): + assert not is_trusted_source_mint(token_mint) + + +@pytest.mark.parametrize( + "token_mint", + [ + "https://mint.exa\tmple", + "https://mint.exa\nmple", + "https://mint.exa\rmple", + ], +) +def test_trusted_mint_matching_rejects_embedded_control_characters( + token_mint: str, +) -> None: + """``urlsplit`` deletes tab/CR/LF before parsing, so such a URL would be + checked as one string and dialled as another.""" + with patch.object(settings, "primary_mint", "https://mint.example"): + with patch.object(settings, "cashu_mints", []): + assert not is_trusted_source_mint(token_mint) + + +@pytest.mark.parametrize( + "token_mint", + ["", " ", "https://", "://mint.example", "mint.example"], +) +def test_trusted_mint_matching_rejects_degenerate_urls(token_mint: str) -> None: + with patch.object(settings, "primary_mint", "https://mint.example"): + with patch.object(settings, "cashu_mints", []): + assert not is_trusted_source_mint(token_mint) + + +def test_unset_primary_mint_never_makes_a_token_trusted() -> None: + """An unset primary mint must not turn an empty token mint into a match.""" + with patch.object(settings, "primary_mint", ""): + with patch.object(settings, "cashu_mints", []): + assert not is_trusted_source_mint("") + assert not is_trusted_source_mint("https://evil.example") + + +def test_trusted_mint_matching_keeps_path_case_sensitive() -> None: + with patch.object(settings, "primary_mint", "https://mint.minibits.cash/Bitcoin"): + with patch.object(settings, "cashu_mints", []): + assert is_trusted_source_mint("https://mint.minibits.cash/Bitcoin/") + assert not is_trusted_source_mint("https://mint.minibits.cash/bitcoin") + + +def test_trusted_mint_matching_rejects_unparseable_port() -> None: + with patch.object(settings, "primary_mint", "https://mint.example"): + with patch.object(settings, "cashu_mints", []): + assert not is_trusted_source_mint("https://mint.example:notaport") + + +def test_trusted_mint_matching_rejects_malformed_url() -> None: + with patch.object(settings, "primary_mint", "https://mint.example"): + with patch.object(settings, "cashu_mints", []): + assert not is_trusted_source_mint("https://[::1/Bitcoin") + assert resolve_trusted_source_mint("https://[::1/Bitcoin") is None + + +def test_resolve_returns_operator_spelling() -> None: + configured = "https://mint.example/Bitcoin" + with patch.object(settings, "primary_mint", configured): + with patch.object(settings, "cashu_mints", []): + assert ( + resolve_trusted_source_mint("HTTPS://MINT.EXAMPLE:443/Bitcoin///") + == configured + ) + + +@pytest.mark.asyncio +async def test_recieve_token_uses_canonical_mint_url() -> None: + variant = PRIMARY.upper() + "///" + get_wallet = AsyncMock(return_value=object()) + redeem = AsyncMock(return_value=(90, "sat", variant)) + with ( + _trusted_mints(), + patch( + "routstr.wallet.deserialize_token_from_string", + return_value=_token(variant), + ), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet._redeem_same_mint", redeem), + ): + amount, unit, mint_url = await recieve_token("cashuAvariant") + + assert (amount, unit, mint_url) == (90, "sat", PRIMARY) + get_wallet.assert_awaited_once_with(PRIMARY, "sat", load=False) + + +def test_classification_has_dedicated_type_and_code() -> None: + classified = classify_redemption_error(UntrustedSourceMintError("x")) + assert classified == ( + "untrusted_mint", + 400, + "Cashu token was issued by a mint this node does not accept", + "cashu_untrusted_source_mint", + ) + + +@pytest.mark.asyncio +async def test_recieve_token_rejects_untrusted_mint_before_mint_contact() -> None: + """The gate runs inside the wallet lock but before ``get_wallet``, so an + untrusted token never reaches the mint over the network.""" + get_wallet = AsyncMock() + with ( + _trusted_mints(), + patch( + "routstr.wallet.deserialize_token_from_string", + return_value=_token(UNTRUSTED), + ), + patch("routstr.wallet.get_wallet", get_wallet), + ): + with pytest.raises(UntrustedSourceMintError): + await recieve_token("cashuAuntrusted") + + get_wallet.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bearer_untrusted_mint_returns_400_with_dedicated_code( + session: AsyncSession, +) -> None: + get_wallet = AsyncMock() + with ( + _trusted_mints(), + patch( + "routstr.auth.deserialize_token_from_string", + return_value=_token(UNTRUSTED), + ), + patch( + "routstr.wallet.deserialize_token_from_string", + return_value=_token(UNTRUSTED), + ), + patch("routstr.wallet.get_wallet", get_wallet), + ): + with pytest.raises(HTTPException) as exc_info: + await validate_bearer_key("cashuAuntrusted", session) + + assert exc_info.value.status_code == 400 + detail = cast(dict[str, dict[str, str]], exc_info.value.detail) + assert detail["error"]["type"] == "untrusted_mint" + assert detail["error"]["code"] == "cashu_untrusted_source_mint" + get_wallet.assert_not_awaited() + + +def test_check_token_balance_rejects_untrusted_mint() -> None: + with ( + _trusted_mints(), + patch( + "routstr.payment.helpers.deserialize_token_from_string", + return_value=_token(UNTRUSTED), + ), + ): + with pytest.raises(HTTPException) as exc_info: + check_token_balance({"x-cashu": "cashuAuntrusted"}, {"model": "m"}, 1) + + assert exc_info.value.status_code == 400 + detail = cast(dict[str, dict[str, str]], exc_info.value.detail) + assert detail["error"]["type"] == "untrusted_mint" + assert detail["error"]["code"] == "cashu_untrusted_source_mint" + + +def test_check_token_balance_accepts_trusted_mints() -> None: + for mint in (PRIMARY, SECONDARY): + with ( + _trusted_mints(), + patch( + "routstr.payment.helpers.deserialize_token_from_string", + return_value=_token(mint), + ), + ): + check_token_balance({"x-cashu": "cashuAtrusted"}, {"model": "m"}, 1) + + +def _http_429(retry_after: str | None) -> httpx.HTTPStatusError: + request = httpx.Request("POST", "http://primary:3338/v1/swap") + headers = {"Retry-After": retry_after} if retry_after else {} + response = httpx.Response(429, request=request, headers=headers) + return httpx.HTTPStatusError("rate limited", request=request, response=response) + + +@pytest.mark.parametrize( + "error", + [_http_429("42"), _http_429(None), MintCooldownError(PRIMARY, 12.4)], +) +def test_rate_limit_asks_to_retry_later(error: Exception) -> None: + assert classify_redemption_error(error) == ( + "mint_rate_limited", + 503, + "Cashu mint is rate-limiting requests; retry later", + "cashu_mint_rate_limited", + ) + + +def test_timeout_has_its_own_code() -> None: + wrapped = SourceMintConnectionError("Issuing Cashu mint is unreachable") + wrapped.__cause__ = httpx.ReadTimeout("read timed out") + assert classify_redemption_error(wrapped) == ( + "mint_timeout", + 503, + "Cashu mint did not respond in time; retry later", + "cashu_mint_timeout", + ) + + +def test_source_mint_unreachable_asks_to_retry() -> None: + wrapped = SourceMintConnectionError("Issuing Cashu mint is unreachable") + wrapped.__cause__ = httpx.ConnectError("refused") + assert classify_redemption_error(wrapped) == ( + "mint_unreachable", + 503, + "The mint that issued this Cashu token is unreachable; retry later", + "cashu_source_mint_unreachable", + ) + + +def test_timeout_wrapped_in_consumed_token_is_not_retryable() -> None: + consumed = TokenConsumedError("credit failed after melt") + consumed.__cause__ = httpx.ReadTimeout("read timed out") + classified = classify_redemption_error(consumed) + assert classified is not None + assert classified[3] == "cashu_token_consumed" + assert not is_mint_timeout(consumed) diff --git a/tests/unit/test_client_app_logging.py b/tests/unit/test_client_app_logging.py new file mode 100644 index 00000000..107f8c2c --- /dev/null +++ b/tests/unit/test_client_app_logging.py @@ -0,0 +1,198 @@ +"""Tests for client-app identification in request logging.""" + +import asyncio +import logging + +import pytest +from fastapi import FastAPI, Request, Response +from fastapi.testclient import TestClient +from starlette.datastructures import Headers + +from routstr.core.logging import ClientAppFilter +from routstr.core.middleware import ( + UNKNOWN_CLIENT_APP, + LoggingMiddleware, + client_app_context, + client_app_from_headers, +) + + +def _record() -> logging.LogRecord: + return logging.LogRecord( + name="routstr.test", + level=logging.INFO, + pathname=__file__, + lineno=1, + msg="test", + args=None, + exc_info=None, + ) + + +@pytest.mark.parametrize( + ("headers", "expected"), + [ + ( + { + "x-title": "Goose", + "http-referer": "https://myapp.example.com", + "user-agent": "python-httpx/0.27", + }, + "Goose", + ), + ( + {"http-referer": "https://myapp.example.com", "user-agent": "curl/8.4.0"}, + "https://myapp.example.com", + ), + ( + {"referer": "https://myapp.example.com", "user-agent": "curl/8.4.0"}, + "https://myapp.example.com", + ), + ({"user-agent": "curl/8.4.0"}, "curl/8.4.0"), + ({}, UNKNOWN_CLIENT_APP), + ({"x-title": " ", "user-agent": "curl/8.4.0"}, "curl/8.4.0"), + ({"x-title": " ", "user-agent": "\t"}, UNKNOWN_CLIENT_APP), + ], + ids=[ + "x-title-wins", + "http-referer", + "referer", + "user-agent-fallback", + "no-identity-headers", + "blank-falls-through", + "all-blank", + ], +) +def test_client_app_from_headers(headers: dict[str, str], expected: str) -> None: + assert client_app_from_headers(Headers(headers)) == expected + + +@pytest.mark.parametrize("header", ["http-referer", "referer"]) +@pytest.mark.parametrize( + ("url", "expected"), + [ + ( + "https://alice:password@app.example:8443/private/chat?token=secret#access_token=secret", + "https://app.example:8443", + ), + ("http://[::1]:3000/chat?key=secret", "http://[::1]:3000"), + ("https://app.example/" + "a" * 200, "https://app.example"), + ("https://[invalid", "curl/8.4.0"), + ("/private/chat?token=secret", "curl/8.4.0"), + ("javascript:secret", "curl/8.4.0"), + ("https:///private", "curl/8.4.0"), + ], +) +def test_referrer_only_identifies_origin(header: str, url: str, expected: str) -> None: + headers = Headers({header: url, "user-agent": "curl/8.4.0"}) + assert client_app_from_headers(headers) == expected + + +def test_value_is_truncated_to_120_chars() -> None: + assert client_app_from_headers(Headers({"x-title": "a" * 500})) == "a" * 120 + + +def test_control_characters_are_stripped() -> None: + headers = Headers({"user-agent": "evil-app\x1b[0m fake INFO line"}) + assert client_app_from_headers(headers) == "evil-app[0m fake INFO line" + + +def test_filter_reads_context_variable() -> None: + token = client_app_context.set("Goose") + try: + record = _record() + assert ClientAppFilter().filter(record) is True + assert record.client_app == "Goose" # type: ignore[attr-defined] + finally: + client_app_context.reset(token) + + +@pytest.mark.parametrize("fail", [False, True]) +async def test_context_is_restored_after_request(fail: bool) -> None: + middleware = LoggingMiddleware(FastAPI()) + request = Request( + { + "type": "http", + "method": "GET", + "path": "/test", + "query_string": b"", + "headers": [], + } + ) + + async def call_next(request: Request) -> Response: + assert client_app_context.get() == UNKNOWN_CLIENT_APP + if fail: + raise RuntimeError("handler failed") + return Response() + + token = client_app_context.set("outer") + try: + if fail: + with pytest.raises(RuntimeError, match="handler failed"): + await middleware.dispatch(request, call_next) + else: + await middleware.dispatch(request, call_next) + assert client_app_context.get() == "outer" + finally: + client_app_context.reset(token) + + +async def test_concurrent_requests_keep_their_own_client_app() -> None: + middleware = LoggingMiddleware(FastAPI()) + ready = asyncio.Event() + apps: list[str] = [] + + async def call_next(request: Request) -> Response: + apps.append(request.headers["x-title"]) + if len(apps) == 2: + ready.set() + await asyncio.wait_for(ready.wait(), timeout=5) + assert client_app_context.get() == request.headers["x-title"] + return Response() + + await asyncio.gather( + *( + middleware.dispatch( + Request( + { + "type": "http", + "method": "GET", + "path": "/test", + "query_string": b"", + "headers": [(b"x-title", app)], + } + ), + call_next, + ) + for app in (b"Goose", b"Pi") + ) + ) + + +def test_filter_defaults_to_unknown_outside_request_context() -> None: + record = _record() + assert ClientAppFilter().filter(record) is True + assert record.client_app == UNKNOWN_CLIENT_APP # type: ignore[attr-defined] + + +def test_handler_logs_carry_client_app(caplog: pytest.LogCaptureFixture) -> None: + app = FastAPI() + handler_logger = logging.getLogger("routstr.test.handler") + + @app.get("/whoami") + async def whoami() -> dict[str, bool]: + handler_logger.warning("something went wrong") + return {"ok": True} + + app.add_middleware(LoggingMiddleware) + + caplog.handler.addFilter(ClientAppFilter()) + handler_logger.addHandler(caplog.handler) + try: + TestClient(app).get("/whoami", headers={"X-Title": "Goose"}) + finally: + handler_logger.removeHandler(caplog.handler) + + record = next(r for r in caplog.records if r.name == "routstr.test.handler") + assert record.client_app == "Goose" # type: ignore[attr-defined] diff --git a/tests/unit/test_completions_billing.py b/tests/unit/test_completions_billing.py new file mode 100644 index 00000000..6a512136 --- /dev/null +++ b/tests/unit/test_completions_billing.py @@ -0,0 +1,384 @@ +import json +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from fastapi.responses import Response +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 +from routstr.auth import ReservationSnapshot, get_reservation_snapshot, pay_for_request +from routstr.core.db import ApiKey, ReservationRelease +from routstr.payment.models import Architecture, Model, Pricing +from routstr.upstream.base import BaseUpstreamProvider + +BALANCE = 100_000 +RESERVED = 5_000 +# 1 msat per prompt token, 2 msat per completion token. +MODEL = Model( + id="glm-test", + name="glm-test", + created=0, + description="", + context_length=64_000, + architecture=Architecture( + modality="text->text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="Other", + instruct_type=None, + ), + pricing=Pricing(prompt=0.001, completion=0.002), + sats_pricing=Pricing(prompt=0.001, completion=0.002), +) +USAGE = {"prompt_tokens": 400, "completion_tokens": 100, "total_tokens": 500} +EXPECTED_MSATS = 400 * 1 + 100 * 2 + +COMPLETION_BODY = { + "model": MODEL.id, + "prompt": "Once upon a time, in a land far away, " * 20, + "max_tokens": 100, +} +CHAT_BODY = {"model": MODEL.id, "messages": [{"role": "user", "content": "hi"}]} + +COMPLETION_JSON = { + "id": "cmpl-1", + "object": "text_completion", + "model": MODEL.id, + "choices": [{"text": " there was", "index": 0, "finish_reason": "stop"}], + "usage": USAGE, +} +COMPLETION_CHUNKS = [ + { + "id": "cmpl-1", + "object": "text_completion", + "model": MODEL.id, + "choices": [{"text": " there", "index": 0, "finish_reason": None}], + }, + { + "id": "cmpl-1", + "object": "text_completion", + "model": MODEL.id, + "choices": [{"text": " was", "index": 0, "finish_reason": "stop"}], + }, +] +USAGE_CHUNK = { + "id": "cmpl-1", + "object": "text_completion", + "model": MODEL.id, + "choices": [], + "usage": USAGE, +} + + +def _sse(chunks: list[dict]) -> bytes: + body = b"".join(b"data: " + json.dumps(c).encode() + b"\n\n" for c in chunks) + return body + b"data: [DONE]\n\n" + + +@pytest.fixture(autouse=True) +def patch_sats_usd_price() -> Any: + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-4): + yield + + +async def _engine() -> AsyncEngine: + engine = create_async_engine("sqlite+aiosqlite://") + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + return engine + + +def _upstream(content: bytes, content_type: str) -> httpx.Response: + return httpx.Response( + 200, + content=content, + headers={"content-type": content_type}, + request=httpx.Request("POST", "http://upstream"), + ) + + +async def _drain(response: Any) -> bytes: + body = b"" + if hasattr(response, "body_iterator"): + async for chunk in response.body_iterator: + body += chunk if isinstance(chunk, bytes) else chunk.encode() + else: + body = response.body + return body + + +async def _forward( + engine: AsyncEngine, + path: str, + body: dict, + upstream: httpx.Response, +) -> tuple[bytes, ReservationSnapshot, AsyncMock]: + """Reserve, forward through the real ``forward_request`` and settle.""" + provider = BaseUpstreamProvider( + base_url="http://upstream", api_key="k", provider_fee=1.0 + ) + request = MagicMock() + request.method = "POST" + request.query_params = {} + send = AsyncMock(return_value=upstream) + + async with AsyncSession(engine, expire_on_commit=False) as session: + key = ApiKey(hashed_key="key", balance=BALANCE) + session.add(key) + await session.commit() + await pay_for_request(key, RESERVED, session) + snapshot = await get_reservation_snapshot(key, session) + + with ( + patch("httpx.AsyncClient.send", send), + patch( + "routstr.upstream.base.create_session", + side_effect=lambda: AsyncSession(engine, expire_on_commit=False), + ), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + auth_module.adjust_payment_for_tokens, + ), + ): + response = await provider.forward_request( + request, + path, + {}, + json.dumps(body).encode(), + key, + RESERVED, + session, + MODEL, + snapshot, + ) + out = await _drain(response) + return out, snapshot, send + + +async def _ledger( + engine: AsyncEngine, snapshot: ReservationSnapshot +) -> tuple[int, int, int, str | None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + key = await session.get(ApiKey, snapshot.key_hash) + record = await session.get(ReservationRelease, snapshot.release_id) + assert key is not None + return ( + key.balance, + key.total_spent, + key.reserved_balance, + record.status if record else None, + ) + + +def _sse_objects(out: bytes) -> list[dict]: + objs = [] + for line in out.split(b"\n"): + if line.startswith(b"data: ") and line[6:].strip() != b"[DONE]": + objs.append(json.loads(line[6:])) + return objs + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "path", + [ + "completions", + "v1/completions", + "v1/completions/", + "openai/v1/completions", + ], +) +async def test_non_streaming_completion_with_usage_is_charged(path: str) -> None: + engine = await _engine() + out, snapshot, _ = await _forward( + engine, + path, + COMPLETION_BODY, + _upstream(json.dumps(COMPLETION_JSON).encode(), "application/json"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert (balance, spent, reserved, status) == ( + BALANCE - EXPECTED_MSATS, + EXPECTED_MSATS, + 0, + "charged", + ) + body = json.loads(out) + assert body["object"] == "text_completion" + assert body["usage"]["cost"]["total_msats"] == EXPECTED_MSATS + await engine.dispose() + + +@pytest.mark.asyncio +async def test_streaming_completion_with_final_usage_is_charged() -> None: + engine = await _engine() + out, snapshot, send = await _forward( + engine, + "v1/completions", + {**COMPLETION_BODY, "stream": True}, + _upstream(_sse([*COMPLETION_CHUNKS, USAGE_CHUNK]), "text/event-stream"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert (balance, spent, reserved, status) == ( + BALANCE - EXPECTED_MSATS, + EXPECTED_MSATS, + 0, + "charged", + ) + + forwarded = json.loads(send.call_args.args[0].content) + assert forwarded["stream_options"] == {"include_usage": True} + + objs = _sse_objects(out) + assert [o["choices"][0]["text"] for o in objs if o["choices"]] == [ + " there", + " was", + ] + assert objs[-1]["object"] == "text_completion" + assert objs[-1]["usage"]["cost"]["total_msats"] == EXPECTED_MSATS + assert out.endswith(b"data: [DONE]\n\n") + await engine.dispose() + + +@pytest.mark.asyncio +async def test_non_streaming_completion_with_empty_usage_is_estimated() -> None: + """Missing usage is estimated from ``prompt`` and ``text``, never free.""" + engine = await _engine() + no_usage = {k: v for k, v in COMPLETION_JSON.items() if k != "id"} + no_usage["usage"] = {} + out, snapshot, _ = await _forward( + engine, + "v1/completions", + COMPLETION_BODY, + _upstream(json.dumps(no_usage).encode(), "application/json"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert 0 < spent <= RESERVED + assert (BALANCE - balance, reserved, status) == (spent, 0, "charged") + body = json.loads(out) + usage = body["usage"] + assert body["id"].startswith("cmpl-") + assert usage["estimated"] is True + assert usage["prompt_tokens"] > 100 + assert usage["completion_tokens"] > 0 + await engine.dispose() + + +@pytest.mark.asyncio +async def test_streaming_completion_without_usage_is_estimated() -> None: + engine = await _engine() + out, snapshot, _ = await _forward( + engine, + "v1/completions", + {**COMPLETION_BODY, "stream": True}, + _upstream(_sse(COMPLETION_CHUNKS), "text/event-stream"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert 0 < spent <= RESERVED + assert (BALANCE - balance, reserved, status) == (spent, 0, "charged") + trailer = _sse_objects(out)[-1] + assert trailer["id"] == "cmpl-1" + assert trailer["object"] == "text_completion" + assert trailer["usage"]["prompt_tokens"] > 100 + assert trailer["usage"]["cost"]["total_msats"] == spent + await engine.dispose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["v1/chat/completions/", "openai/v1/chat/completions"]) +async def test_chat_completion_aliases_are_charged(path: str) -> None: + engine = await _engine() + chat_json = { + "id": "chatcmpl-1", + "object": "chat.completion", + "model": MODEL.id, + "choices": [ + { + "message": {"role": "assistant", "content": "hi"}, + "index": 0, + "finish_reason": "stop", + } + ], + "usage": USAGE, + } + out, snapshot, _ = await _forward( + engine, + path, + CHAT_BODY, + _upstream(json.dumps(chat_json).encode(), "application/json"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert (balance, spent, reserved, status) == ( + BALANCE - EXPECTED_MSATS, + EXPECTED_MSATS, + 0, + "charged", + ) + assert json.loads(out)["usage"]["cost"]["total_msats"] == EXPECTED_MSATS + await engine.dispose() + + +def test_stream_usage_option_is_scoped_to_completion_endpoints() -> None: + provider = BaseUpstreamProvider(base_url="http://upstream", api_key="k") + body = json.dumps({"prompt": "draw this", "stream": True}).encode() + + assert provider.prepare_request_body(body, MODEL) == body + + +async def _forward_x_cashu( + path: str, body: dict, upstream: httpx.Response +) -> tuple[AsyncMock, AsyncMock]: + """Run ``forward_x_cashu_request`` with the settlement handler stubbed out.""" + provider = BaseUpstreamProvider( + base_url="http://upstream", api_key="k", provider_fee=1.0 + ) + request = MagicMock() + request.method = "POST" + request.query_params = {} + request.state.request_id = "req-1" + request.body = AsyncMock(return_value=json.dumps(body).encode()) + send = AsyncMock(return_value=upstream) + settle = AsyncMock(return_value=Response(content=b"{}", status_code=200)) + + with ( + patch("httpx.AsyncClient.send", send), + patch.object(provider, "handle_x_cashu_chat_completion", settle), + ): + await provider.forward_x_cashu_request( + request, path, {}, 10, "sat", RESERVED, MODEL + ) + return settle, send + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["completions", "v1/completions", "v1/completions/"]) +async def test_x_cashu_legacy_completion_is_settled(path: str) -> None: + """Legacy completions must reach refund settlement, not raw passthrough.""" + settle, _ = await _forward_x_cashu( + path, + COMPLETION_BODY, + _upstream(json.dumps(COMPLETION_JSON).encode(), "application/json"), + ) + + assert settle.await_count == 1 + + +@pytest.mark.asyncio +async def test_x_cashu_streaming_completion_requests_usage() -> None: + _, send = await _forward_x_cashu( + "v1/completions", + {**COMPLETION_BODY, "stream": True}, + _upstream(_sse([*COMPLETION_CHUNKS, USAGE_CHUNK]), "text/event-stream"), + ) + + forwarded = json.loads(send.call_args.args[0].content) + assert forwarded["stream_options"] == {"include_usage": True} diff --git a/tests/unit/test_core_exceptions.py b/tests/unit/test_core_exceptions.py index 3d9469df..5f4cc577 100644 --- a/tests/unit/test_core_exceptions.py +++ b/tests/unit/test_core_exceptions.py @@ -1,4 +1,5 @@ import json +from unittest.mock import patch import pytest from fastapi import HTTPException @@ -34,11 +35,14 @@ async def test_structured_http_error_uses_standard_error_envelope() -> None: "details": {"mint": "https://mint.example"}, } - response = await http_exception_handler( - request, - HTTPException(status_code=503, detail={"error": error}), - ) + with patch("routstr.core.exceptions.logger") as logger: + response = await http_exception_handler( + request, + HTTPException(status_code=503, detail={"error": error}), + ) + logger.warning.assert_called_once() + logger.error.assert_not_called() assert response.status_code == 503 assert json.loads(response.body) == { "detail": {"error": error}, diff --git a/tests/unit/test_cost_error_after_delivery.py b/tests/unit/test_cost_error_after_delivery.py new file mode 100644 index 00000000..5540836d --- /dev/null +++ b/tests/unit/test_cost_error_after_delivery.py @@ -0,0 +1,47 @@ +"""A pricing failure after the upstream served content must not raise a 400.""" + +from unittest.mock import AsyncMock, patch + +import pytest + +from routstr.auth import ReservationSnapshot, adjust_payment_for_tokens +from routstr.core.db import ApiKey +from routstr.payment import cost_calculation + + +@pytest.mark.asyncio +async def test_cost_data_error_releases_without_raising() -> None: + key_hash = "a" * 64 + reservation = ReservationSnapshot( + release_id="rel-1", + key_hash=key_hash, + billing_key_hash=key_hash, + reserved_msats=7_000, + ) + with ( + patch.object( + cost_calculation, + "_get_pricing_rates", + side_effect=ValueError("no pricing for model"), + ), + patch("routstr.auth._validate_reservation_snapshot", new=AsyncMock()), + patch("routstr.auth._stop_reservation_heartbeat", new=AsyncMock()), + patch( + "routstr.auth._claim_reservation_for_charge", + new=AsyncMock(return_value=True), + ), + patch( + "routstr.auth._charge_reservation_rows", new=AsyncMock(return_value=True) + ), + patch("routstr.auth.accumulate_routstr_fee", new=AsyncMock()), + ): + cost = await adjust_payment_for_tokens( + ApiKey(hashed_key=key_hash), + {"model": "gpt-4o", "usage": {"prompt_tokens": 10, "completion_tokens": 5}}, + session=AsyncMock(), + deducted_max_cost=7_000, + reservation_snapshot=reservation, + ) + + assert cost["total_msats"] == 0 + assert cost["charged_msats"] == 0 diff --git a/tests/unit/test_cost_response_metadata.py b/tests/unit/test_cost_response_metadata.py index beaf4005..3f06fac7 100644 --- a/tests/unit/test_cost_response_metadata.py +++ b/tests/unit/test_cost_response_metadata.py @@ -58,6 +58,7 @@ def _assert_cost_contract(response: Any) -> None: "input_msats": 1_200, "output_msats": 300, "total_msats": 1_500, + "charged_msats": 1_500, "total_usd": 0.0001, "cache_read_input_tokens": 8, "cache_creation_input_tokens": 2, @@ -69,6 +70,36 @@ def _assert_cost_contract(response: Any) -> None: assert response.headers["X-Routstr-Output-Cost-Msats"] == "300" +@pytest.mark.asyncio +async def test_duplicate_finalization_publishes_zero_debit_and_computed_usage() -> None: + provider = _provider() + duplicate_cost = {**COST_DATA, "charged_msats": 0} + with patch( + "routstr.upstream.base.adjust_payment_for_tokens", + new=AsyncMock(return_value=duplicate_cost), + ): + response = await provider.handle_non_streaming_chat_completion( + _upstream_response( + { + "model": "test-model", + "usage": {"prompt_tokens": 10, "completion_tokens": 3}, + } + ), + _key(), + _session(), + deducted_max_cost=10_000, + ) + + body = json.loads(response.body) + assert response.headers["X-Routstr-Cost-Msats"] == "0" + assert response.headers["X-Routstr-Computed-Cost-Msats"] == "1500" + assert body["usage"]["cost"]["total_msats"] == 0 + assert body["usage"]["cost"]["charged_msats"] == 0 + assert body["usage"]["cost"]["computed_msats"] == 1_500 + assert body["cost"]["total_msats"] == 0 + assert body["cost"]["computed_msats"] == 1_500 + + @pytest.mark.asyncio async def test_balance_chat_completion_uses_shared_cost_contract() -> None: provider = _provider() diff --git a/tests/unit/test_count_tokens_local.py b/tests/unit/test_count_tokens_local.py index 6435948b..23bf379c 100644 --- a/tests/unit/test_count_tokens_local.py +++ b/tests/unit/test_count_tokens_local.py @@ -13,7 +13,7 @@ from unittest.mock import patch from routstr.payment.models import Architecture, Model, Pricing from routstr.upstream import count_tokens as count_tokens_module -from routstr.upstream.count_tokens import count_tokens_locally +from routstr.upstream.count_tokens import MissingUsageEstimator, count_tokens_locally def _make_model(model_id: str = "anthropic/claude-3-5-sonnet") -> Model: @@ -154,6 +154,103 @@ def test_supports_anthropic_system_block_list() -> None: assert payload["input_tokens"] > 0 +def test_missing_usage_estimator_prices_request_and_streamed_output() -> None: + model = _make_model() + request_body = _body( + { + "model": model.id, + "messages": [{"role": "user", "content": "price this prompt"}], + } + ) + + with ( + patch.object(count_tokens_module, "_count_with_litellm", return_value=17), + patch.object(count_tokens_module, "_count_text_with_litellm", return_value=5), + ): + estimator = MissingUsageEstimator(request_body, model) + estimator.observe( + { + "model": "provider/model", + "choices": [{"delta": {"content": "estimated output"}}], + } + ) + response = estimator.response_data("provider/model") + + assert response == { + "model": "provider/model", + "usage": { + "input_tokens": 17, + "output_tokens": 5, + "total_tokens": 22, + "estimated": True, + }, + } + + +def test_missing_usage_estimator_skips_responses_api_done_events() -> None: + estimator = MissingUsageEstimator(b"{}", None) + estimator.observe({"type": "response.output_text.delta", "delta": "streamed"}) + estimator.observe({"type": "response.output_text.done", "text": "streamed"}) + estimator.observe( + { + "type": "response.content_part.done", + "part": {"type": "output_text", "text": "streamed"}, + } + ) + + assert estimator.output_text == "streamed" + + +def test_missing_usage_estimator_openai_dialect() -> None: + model = _make_model() + request_body = _body( + { + "model": model.id, + "messages": [{"role": "user", "content": "price this prompt"}], + } + ) + + with ( + patch.object(count_tokens_module, "_count_with_litellm", return_value=17), + patch.object(count_tokens_module, "_count_text_with_litellm", return_value=5), + ): + estimator = MissingUsageEstimator(request_body, model) + estimator.observe({"choices": [{"delta": {"content": "estimated output"}}]}) + response = estimator.openai_response_data("provider/model") + + assert response == { + "model": "provider/model", + "usage": { + "prompt_tokens": 17, + "completion_tokens": 5, + "total_tokens": 22, + "estimated": True, + }, + } + + +def test_missing_usage_estimator_counts_legacy_token_prompt() -> None: + model = _make_model() + request_body = _body({"model": model.id, "prompt": [[1, 2], [3, 4, 5]]}) + + usage = MissingUsageEstimator(request_body, model).openai_response_data()["usage"] + + assert usage["prompt_tokens"] >= 5 + + +def test_missing_usage_estimator_does_not_count_response_metadata() -> None: + estimator = MissingUsageEstimator(b"{}", None) + estimator.observe( + { + "id": "chatcmpl-this-is-not-generated-text", + "model": "also-not-generated-text", + "choices": [{"delta": {"role": "assistant"}}], + } + ) + + assert estimator.output_text == "" + + def test_uses_forwarded_model_id_when_present() -> None: model = _make_model("anthropic/claude-3-5-sonnet") model.forwarded_model_id = "claude-3-5-sonnet-20241022" @@ -175,3 +272,39 @@ def test_uses_forwarded_model_id_when_present() -> None: assert captured["model"] == "claude-3-5-sonnet-20241022" assert _read_payload(response)["input_tokens"] == 7 + + +def test_responses_instructions_are_counted_as_system_text() -> None: + body = {"model": "gpt-4o", "input": "Hi", "instructions": "Be concise."} + with patch.object( + count_tokens_module.litellm, "token_counter", return_value=12 + ) as counter: + usage = MissingUsageEstimator(_body(body), None).response_data()["usage"] + + assert usage["input_tokens"] == 12 + counter.assert_called_once_with( + model="gpt-4o", + messages=[ + {"role": "system", "content": "Be concise."}, + {"role": "user", "content": "Hi"}, + ], + tools=None, + ) + + +def test_responses_tool_results_use_fallback_instead_of_empty_messages() -> None: + body = { + "model": "gpt-4o", + "input": [ + { + "type": "function_call_output", + "call_id": "call_1", + "output": "result " * 100, + } + ], + } + with patch.object(count_tokens_module.litellm, "token_counter") as counter: + usage = MissingUsageEstimator(_body(body), None).response_data()["usage"] + + counter.assert_not_called() + assert usage["input_tokens"] > 100 diff --git a/tests/unit/test_coverage_payment_helpers.py b/tests/unit/test_coverage_payment_helpers.py index b6ac57ed..415b67f3 100644 --- a/tests/unit/test_coverage_payment_helpers.py +++ b/tests/unit/test_coverage_payment_helpers.py @@ -8,6 +8,8 @@ from unittest.mock import Mock, patch import pytest +from routstr.core.settings import settings + # --------------------------------------------------------------------------- # check_token_balance # --------------------------------------------------------------------------- @@ -22,6 +24,7 @@ async def test_check_token_balance_x_cashu_present() -> None: with patch("routstr.payment.helpers.deserialize_token_from_string") as mock_deser: mock_token = Mock() + mock_token.mint = settings.primary_mint mock_token.amount = 50000 mock_token.unit = "sat" mock_deser.return_value = mock_token @@ -62,6 +65,7 @@ async def test_check_token_balance_insufficient_raises() -> None: with patch("routstr.payment.helpers.deserialize_token_from_string") as mock_deser: mock_token = Mock() + mock_token.mint = settings.primary_mint mock_token.amount = 100 # 100 sat mock_token.unit = "sat" mock_deser.return_value = mock_token diff --git a/tests/unit/test_db_pool_config.py b/tests/unit/test_db_pool_config.py index 8f5d0367..bc2a8dd9 100644 --- a/tests/unit/test_db_pool_config.py +++ b/tests/unit/test_db_pool_config.py @@ -56,9 +56,40 @@ def test_non_sqlite_backend_enables_pre_ping_automatically( assert created is fake_engine assert factory.call_args.kwargs["pool_pre_ping"] is True + assert "timeout" not in factory.call_args.kwargs["connect_args"] assert listen.call_count == 2 +def test_file_sqlite_sets_busy_timeout_connect_arg( + monkeypatch: pytest.MonkeyPatch, tmp_path: object +) -> None: + monkeypatch.setattr(settings, "database_busy_timeout", 42.0) + fake_engine = MagicMock() + + with ( + patch.object(db, "create_async_engine", return_value=fake_engine) as factory, + patch.object(db.event, "listen"), + ): + create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/busy.db") + + assert factory.call_args.kwargs["connect_args"]["timeout"] == 42.0 + + +def test_memory_sqlite_omits_busy_timeout_connect_arg( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "database_busy_timeout", 42.0) + fake_engine = MagicMock() + + with ( + patch.object(db, "create_async_engine", return_value=fake_engine) as factory, + patch.object(db.event, "listen"), + ): + create_db_engine("sqlite+aiosqlite://") + + assert "timeout" not in factory.call_args.kwargs["connect_args"] + + @pytest.mark.asyncio async def test_every_created_engine_warns_for_long_checkouts( monkeypatch: pytest.MonkeyPatch, tmp_path: object diff --git a/tests/unit/test_ehbp_finalize_payment.py b/tests/unit/test_ehbp_finalize_payment.py index 08cb12ac..caf0a5aa 100644 --- a/tests/unit/test_ehbp_finalize_payment.py +++ b/tests/unit/test_ehbp_finalize_payment.py @@ -1,6 +1,8 @@ from __future__ import annotations -from typing import Any, AsyncGenerator +import logging +from contextlib import contextmanager +from typing import Any, AsyncGenerator, Iterator from unittest.mock import AsyncMock, MagicMock import pytest @@ -9,9 +11,11 @@ from sqlalchemy.pool import StaticPool from sqlmodel import SQLModel, select from sqlmodel.ext.asyncio.session import AsyncSession +import routstr.auth as auth_module from routstr.auth import get_reservation_snapshot, pay_for_request from routstr.core.db import ApiKey, ReservationRelease from routstr.upstream.ehbp import ( + _inject_cost_response_headers, finalize_ehbp_actual_cost_payment, finalize_ehbp_max_cost_payment, ) @@ -26,7 +30,9 @@ def _make_engine() -> AsyncEngine: @pytest.fixture -async def session(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[AsyncSession, None]: +async def session( + monkeypatch: pytest.MonkeyPatch, +) -> AsyncGenerator[AsyncSession, None]: monkeypatch.setattr("routstr.upstream.ehbp.ROUTSTR_FEE_PERCENT", 0) engine = _make_engine() async with engine.begin() as conn: @@ -35,6 +41,8 @@ async def session(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[AsyncSessio try: yield db_session finally: + for release_id in list(auth_module._reservation_heartbeats): + await auth_module._stop_reservation_heartbeat(release_id) await db_session.close() await engine.dispose() @@ -54,9 +62,7 @@ def _fail_nth_api_key_update( original_exec = session.exec api_key_updates = 0 - async def exec_with_failure( - statement: Any, *args: Any, **kwargs: Any - ) -> Any: + async def exec_with_failure(statement: Any, *args: Any, **kwargs: Any) -> Any: nonlocal api_key_updates table = getattr(statement, "table", None) if getattr(table, "name", None) == "api_keys": @@ -78,7 +84,7 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve await pay_for_request(key, 3_000, session) reservation = await get_reservation_snapshot(key, session) - await finalize_ehbp_actual_cost_payment( + charged = await finalize_ehbp_actual_cost_payment( key, session, reserved_cost_for_model=3_000, @@ -93,6 +99,7 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve reservation_snapshot=reservation, ) + assert charged == 1_200 updated = await _api_key(session, "ehbp-actual") assert updated is not None assert updated.balance == 8_800 @@ -101,48 +108,144 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve assert updated.total_spent == 1_200 +@contextmanager +def _capture_payments_logs() -> Iterator[list[logging.LogRecord]]: + """Collect ``routstr.payments`` records for the duration of the block. + + ``setup_logging()`` sets ``propagate=False`` on the ``routstr`` logger, so + pytest's ``caplog`` (attached at the root) never sees these records; a + handler on the payments logger itself does. + """ + records: list[logging.LogRecord] = [] + + class _RecordingHandler(logging.Handler): + def emit(self, record: logging.LogRecord) -> None: + records.append(record) + + payments_logger = logging.getLogger("routstr.payments") + handler = _RecordingHandler(level=logging.INFO) + previous_level = payments_logger.level + payments_logger.addHandler(handler) + payments_logger.setLevel(logging.INFO) + try: + yield records + finally: + payments_logger.removeHandler(handler) + payments_logger.setLevel(previous_level) + + @pytest.mark.asyncio -async def test_finalize_max_cost_payment_updates_parent_and_child_spend( +async def test_finalize_actual_cost_payment_logs_cache_tokens( session: AsyncSession, ) -> None: - parent = ApiKey(hashed_key="ehbp-parent", balance=10_000) - child = ApiKey( - hashed_key="ehbp-child", balance=0, parent_key_hash="ehbp-parent" - ) - session.add(parent) - session.add(child) + """The FINALIZE event carries the cache splits, not just input/output.""" + key = ApiKey(hashed_key="ehbp-cache-logging", balance=10_000) + session.add(key) await session.commit() - await pay_for_request(child, 3_000, session) - reservation = await get_reservation_snapshot(child, session) + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) - await finalize_ehbp_max_cost_payment( - child, + with _capture_payments_logs() as records: + charged = await finalize_ehbp_actual_cost_payment( + key, + session, + reserved_cost_for_model=3_000, + model_id="tinfoil/glm-5-2", + cost_info={ + "total_msats": 1_200, + "input_tokens": 5, + "output_tokens": 20, + "input_msats": 500, + "output_msats": 700, + "cache_read_input_tokens": 64, + "cache_creation_input_tokens": 0, + "cache_read_msats": 12, + "cache_creation_msats": 0, + }, + reservation_snapshot=reservation, + ) + + assert charged == 1_200 + finalize_records = [ + record for record in records if record.getMessage() == "FINALIZE" + ] + assert len(finalize_records) == 1 + record = finalize_records[0] + # finalize_type/input_tokens/... are attached via logging's extra= payload. + assert record.finalize_type == "ehbp_usage" # type: ignore[attr-defined] + assert record.input_tokens == 5 # type: ignore[attr-defined] + assert record.output_tokens == 20 # type: ignore[attr-defined] + assert record.cache_read_input_tokens == 64 # type: ignore[attr-defined] + assert record.cache_creation_input_tokens == 0 # type: ignore[attr-defined] + assert record.cache_read_msats == 12 # type: ignore[attr-defined] + assert record.cache_creation_msats == 0 # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_finalize_actual_cost_payment_logs_zero_cache_when_absent( + session: AsyncSession, +) -> None: + """Providers that report no cache split still emit a stable key set.""" + key = ApiKey(hashed_key="ehbp-no-cache-logging", balance=10_000) + session.add(key) + await session.commit() + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) + + with _capture_payments_logs() as records: + await finalize_ehbp_actual_cost_payment( + key, + session, + reserved_cost_for_model=3_000, + model_id="tinfoil/glm-5-2", + cost_info={ + "total_msats": 1_200, + "input_tokens": 10, + "output_tokens": 20, + }, + reservation_snapshot=reservation, + ) + + record = next(record for record in records if record.getMessage() == "FINALIZE") + assert record.cache_read_input_tokens == 0 # type: ignore[attr-defined] + assert record.cache_creation_input_tokens == 0 # type: ignore[attr-defined] + assert record.cache_read_msats == 0 # type: ignore[attr-defined] + assert record.cache_creation_msats == 0 # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_unmeasured_ehbp_releases_reservation( + session: AsyncSession, +) -> None: + key = ApiKey(hashed_key="ehbp-key", balance=10_000) + session.add(key) + await session.commit() + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) + + charged = await finalize_ehbp_max_cost_payment( + key, session, max_cost_for_model=3_000, model_id="tinfoil/model", reservation_snapshot=reservation, ) - updated_parent = await _api_key(session, "ehbp-parent") - updated_child = await _api_key(session, "ehbp-child") - assert updated_parent is not None - assert updated_child is not None - assert updated_parent.balance == 7_000 - assert updated_parent.reserved_balance == 0 - assert updated_parent.reserved_at is None - assert updated_parent.total_spent == 3_000 - assert updated_child.balance == 0 - assert updated_child.reserved_balance == 0 - assert updated_child.reserved_at is None - assert updated_child.total_spent == 3_000 + assert charged == 0 + updated = await _api_key(session, "ehbp-key") + assert updated is not None + assert updated.balance == 10_000 + assert updated.reserved_balance == 0 + assert updated.reserved_at is None + assert updated.total_spent == 0 @pytest.mark.asyncio -async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matches_no_rows( +async def test_finalize_actual_cost_payment_rolls_back_when_billing_key_update_matches_no_rows( session: AsyncSession, monkeypatch: pytest.MonkeyPatch, ) -> None: - key = ApiKey(hashed_key="ehbp-missing-parent", balance=10_000) + key = ApiKey(hashed_key="ehbp-failed-billing-update", balance=10_000) session.add(key) await session.commit() await pay_for_request(key, 3_000, session) @@ -151,7 +254,7 @@ async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matche rollback_spy = AsyncMock(wraps=session.rollback) monkeypatch.setattr(session, "rollback", rollback_spy) - await finalize_ehbp_actual_cost_payment( + charged = await finalize_ehbp_actual_cost_payment( key, session, reserved_cost_for_model=3_000, @@ -160,49 +263,64 @@ async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matche reservation_snapshot=reservation, ) + assert charged == 0 rollback_spy.assert_awaited_once() - updated = await _api_key(session, "ehbp-missing-parent") + updated = await _api_key(session, "ehbp-failed-billing-update") assert updated is not None assert updated.balance == 10_000 - assert updated.reserved_balance == 3_000 + assert updated.reserved_balance == 0 assert updated.total_spent == 0 release = await session.get(ReservationRelease, reservation.release_id) assert release is not None - assert release.status == "active" + assert release.status == "released" + assert reservation.release_id not in auth_module._reservation_heartbeats @pytest.mark.asyncio -async def test_finalize_max_cost_payment_rolls_back_parent_when_child_update_matches_no_rows( +async def test_unmeasured_ehbp_release_is_safe_when_charge_update_would_fail( session: AsyncSession, monkeypatch: pytest.MonkeyPatch, ) -> None: - parent = ApiKey(hashed_key="ehbp-rollback-parent", balance=10_000) - child = ApiKey( - hashed_key="ehbp-missing-child", - balance=0, - parent_key_hash="ehbp-rollback-parent", - ) - session.add(parent) - session.add(child) + key = ApiKey(hashed_key="ehbp-rollback-key", balance=10_000) + session.add(key) await session.commit() - await pay_for_request(child, 3_000, session) - reservation = await get_reservation_snapshot(child, session) - _fail_nth_api_key_update(session, monkeypatch, target_update=2) + await pay_for_request(key, 3_000, session) + reservation = await get_reservation_snapshot(key, session) + _fail_nth_api_key_update(session, monkeypatch, target_update=1) - await finalize_ehbp_max_cost_payment( - child, + charged = await finalize_ehbp_max_cost_payment( + key, session, max_cost_for_model=3_000, model_id="tinfoil/model", reservation_snapshot=reservation, ) - updated_parent = await _api_key(session, "ehbp-rollback-parent") - assert updated_parent is not None - assert updated_parent.balance == 10_000 - assert updated_parent.reserved_balance == 3_000 - assert updated_parent.total_spent == 0 - updated_child = await _api_key(session, "ehbp-missing-child") - assert updated_child is not None - assert updated_child.reserved_balance == 3_000 - assert updated_child.total_spent == 0 + assert charged == 0 + updated = await _api_key(session, "ehbp-rollback-key") + assert updated is not None + assert updated.balance == 10_000 + # The injected partial-update failure rolls aggregate subtraction back; + # terminal fencing prevents a charge or retry from consuming those funds. + assert updated.reserved_balance == 3_000 + assert updated.total_spent == 0 + release = await session.get(ReservationRelease, reservation.release_id) + assert release is not None and release.status == "released" + assert reservation.release_id not in auth_module._reservation_heartbeats + + +def test_zero_debit_ehbp_headers_preserve_computed_cost() -> None: + headers: dict[str, str] = {} + + _inject_cost_response_headers( + headers, + { + "total_msats": 0, + "computed_msats": 1_500, + "input_msats": 1_200, + "output_msats": 300, + }, + ) + + assert headers["X-Routstr-Cost-Msats"] == "0" + assert headers["X-Routstr-Computed-Cost-Msats"] == "1500" diff --git a/tests/unit/test_ehbp_timeout.py b/tests/unit/test_ehbp_timeout.py new file mode 100644 index 00000000..c9e749cb --- /dev/null +++ b/tests/unit/test_ehbp_timeout.py @@ -0,0 +1,131 @@ +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from routstr.core.exceptions import EhbpTimeoutError, UpstreamError +from routstr.upstream import ehbp as ehbp_module + +# --------------------------------------------------------------------------- +# forward_ehbp_x_cashu_request — timeout fails closed with a refund + 504 +# --------------------------------------------------------------------------- + + +async def _request() -> MagicMock: + request = MagicMock() + request.state.request_id = "req-123" + request.method = "POST" + request.query_params = {} + request.headers = {} + request.body = AsyncMock(return_value=b"opaque") + return request + + +def _ehbp_upstream_mocks() -> tuple[MagicMock, MagicMock]: + """Upstream and model mocks sufficient to reach the forwarding call.""" + profile = MagicMock() + profile.client_target_url_header = None + profile.allow_client_target_override = False + profile.proxy_only_headers = frozenset() + profile.usage_response_header = None + + target = MagicMock() + target.url = "https://inference.tinfoil.sh/v1/chat/completions" + target.headers = {} + target.profile = None + + upstream = MagicMock() + upstream.prepare_headers.return_value = {} + upstream.get_ehbp_forwarding_target.return_value = target + upstream.get_confidential_inference_profile.return_value = profile + upstream.prepare_params.return_value = {} + + model_obj = MagicMock() + model_obj.id = "tinfoil-kimi-k2-6" + model_obj.forwarded_model_id = "kimi-k2-6" + return upstream, model_obj + + +@pytest.mark.asyncio +async def test_x_cashu_timeout_refunds_and_returns_504( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + ehbp_module, + "recieve_token", + AsyncMock(return_value=(1000, "msat", None)), + ) + monkeypatch.setattr( + ehbp_module, "store_cashu_transaction", AsyncMock(return_value=None) + ) + send_cashu_refund_mock = AsyncMock(return_value="refund-token") + monkeypatch.setattr(ehbp_module, "send_cashu_refund", send_cashu_refund_mock) + monkeypatch.setattr( + ehbp_module, + "forward_with_trailer", + AsyncMock(side_effect=EhbpTimeoutError("EHBP upstream timed out")), + ) + + upstream, model_obj = _ehbp_upstream_mocks() + + response = await ehbp_module.forward_ehbp_x_cashu_request( + request=await _request(), + x_cashu_token="cashu-token", + path="v1/chat/completions", + max_cost_for_model=5000, + model_obj=model_obj, + upstream=upstream, + ) + + assert response.status_code == 504 + assert response.headers["X-Cashu"] == "refund-token" + 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 +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_bearer_timeout_propagates_504( + 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. + """ + monkeypatch.setattr( + ehbp_module, + "forward_with_trailer", + AsyncMock( + side_effect=EhbpTimeoutError( + "EHBP upstream inference.tinfoil.sh timed out after 60s connecting" + ) + ), + ) + upstream, model_obj = _ehbp_upstream_mocks() + key = MagicMock() + key.hashed_key = "abcdef1234567890" + + with pytest.raises(EhbpTimeoutError) 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 == 504 + assert exc_info.value.code == "UPSTREAM_TIMEOUT" + 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 9e1dfaf2..c6fc3073 100644 --- a/tests/unit/test_fee_payout_crash_safety.py +++ b/tests/unit/test_fee_payout_crash_safety.py @@ -11,6 +11,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession from routstr import wallet from routstr.core import db +from routstr.payment.lnurl import LNURLError class _SessionContext: @@ -525,6 +526,47 @@ async def test_fee_payout_keeps_legacy_checkpoint_without_quote_locked() -> None critical.assert_called_once() +@pytest.mark.asyncio +async def test_fee_payout_failure_before_quote_is_not_reported_as_unknown() -> None: + session = Mock() + fee = SimpleNamespace( + accumulated_msats=1_061_000, + payout_in_progress_msats=0, + payout_started_at=None, + ) + reset = AsyncMock() + + with ( + patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1), + patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1), + patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"), + patch( + "routstr.wallet.asyncio.sleep", + AsyncMock(side_effect=[None, asyncio.CancelledError()]), + ), + patch( + "routstr.wallet.db.create_session", return_value=_session_context(session) + ), + patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), + patch("routstr.wallet.db.reset_routstr_fee", reset), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())), + patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]), + patch( + "routstr.wallet.raw_send_to_lnurl", + side_effect=LNURLError("Cashu melt fees leave no payable LNURL amount"), + ), + patch("routstr.wallet.logger.error") as error, + patch("routstr.wallet.logger.critical") as critical, + ): + with pytest.raises(asyncio.CancelledError): + await wallet.periodic_routstr_fee_payout() + + reset.assert_not_awaited() + critical.assert_not_called() + error.assert_called_once() + assert error.call_args.args[0] == "Routstr fee payout failed before melt dispatch" + + @pytest.mark.asyncio async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> None: session = Mock() diff --git a/tests/unit/test_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index ceeca1e4..12e57fd4 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -136,19 +136,20 @@ async def test_supported_mint_units_come_from_active_keysets() -> None: msat = MagicMock(active=False, unit="msat") usd = MagicMock(active=True) usd.unit.name = "usd" - wallet = MagicMock() - wallet._get_keysets = AsyncMock(return_value=[usd, msat, sat]) + wallet = MagicMock(url="http://mint:3338", db=MagicMock()) + get_keysets = AsyncMock(return_value=[usd, msat, sat]) with ( patch.object(settings, "primary_mint_unit", "sat"), patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.wallet.get_cashu_keysets", get_keysets), ): units = await _get_supported_mint_units("http://mint:3338") cached_units = await _get_supported_mint_units("http://mint:3338") assert units == ["sat", "usd"] assert cached_units == units - wallet._get_keysets.assert_awaited_once() + get_keysets.assert_awaited_once_with(mint_url=wallet.url, db=wallet.db) @pytest.mark.asyncio @@ -339,7 +340,7 @@ async def test_slow_mints_do_not_exhaust_a_single_connection_pool( f"sqlite+aiosqlite:///{tmp_path / 'pool-pressure.db'}", pool_size=1, max_overflow=0, - pool_timeout=0.2, + pool_timeout=0.5, ) async with engine.begin() as connection: await connection.run_sync(SQLModel.metadata.create_all) @@ -350,7 +351,7 @@ async def test_slow_mints_do_not_exhaust_a_single_connection_pool( yield session async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def] - await asyncio.sleep(0.3) + await asyncio.sleep(1.0) return proofs try: diff --git a/tests/unit/test_image_url_fetch_guard.py b/tests/unit/test_image_url_fetch_guard.py new file mode 100644 index 00000000..4cca987e --- /dev/null +++ b/tests/unit/test_image_url_fetch_guard.py @@ -0,0 +1,167 @@ +"""Guards for the pre-auth image URL fetch used by cost estimation.""" + +import socket +import threading +from http.server import BaseHTTPRequestHandler, HTTPServer +from typing import Any, Callable, Iterator + +import pytest + +from routstr.payment import helpers +from routstr.payment.helpers import ( + IMAGE_FETCH_MAX_BYTES, + IMAGE_FETCH_MAX_PER_REQUEST, + _fetch_image_from_url, + _is_blocked_address, + _validated_fetch_target, + estimate_image_tokens_in_messages, +) + +REQUESTED_PATHS: list[str] = [] + + +class _Loop: + def __init__(self, getaddrinfo: Callable[..., Any]) -> None: + self.getaddrinfo = getaddrinfo + + +class _Sink(BaseHTTPRequestHandler): + def do_GET(self) -> None: # noqa: N802 + REQUESTED_PATHS.append(self.path) + body = b"x" * (IMAGE_FETCH_MAX_BYTES * 4) + self.send_response(200) + self.send_header("Content-Type", "image/png") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + try: + self.wfile.write(body) + except BrokenPipeError: + # The client stops reading once the byte cap is reached. + pass + + def log_message(self, *args: object) -> None: + pass + + +@pytest.fixture +def sink() -> Iterator[str]: + REQUESTED_PATHS.clear() + server = HTTPServer(("127.0.0.1", 0), _Sink) + threading.Thread(target=server.serve_forever, daemon=True).start() + try: + yield f"http://127.0.0.1:{server.server_address[1]}" + finally: + server.shutdown() + server.server_close() + + +@pytest.mark.asyncio +async def test_loopback_url_is_not_fetched(sink: str) -> None: + assert await _fetch_image_from_url(f"{sink}/internal") is None + assert REQUESTED_PATHS == [] + + +@pytest.mark.asyncio +async def test_link_local_metadata_url_is_not_fetched() -> None: + assert await _fetch_image_from_url("http://169.254.169.254/latest/meta-data") is None + + +@pytest.mark.asyncio +async def test_non_http_scheme_is_rejected() -> None: + assert await _fetch_image_from_url("file:///etc/passwd") is None + + +@pytest.mark.parametrize( + "address", + [ + "127.0.0.1", + "10.0.0.1", + "169.254.169.254", + "100.64.0.1", # CGNAT: reachable inside many hosting networks + "192.0.0.1", + "224.0.0.1", + "::1", + "::ffff:127.0.0.1", + "2002:7f00:1::", # 6to4 wrapping 127.0.0.1 + ], +) +def test_non_global_addresses_are_blocked(address: str) -> None: + assert _is_blocked_address(address) is True + + +@pytest.mark.parametrize("address", ["8.8.8.8", "2001:4860:4860::8888"]) +def test_global_addresses_are_allowed(address: str) -> None: + assert _is_blocked_address(address) is False + + +@pytest.mark.asyncio +async def test_http_target_is_pinned_to_validated_address( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def fake_getaddrinfo(*args: object, **kwargs: object) -> list[tuple]: + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 80))] + + monkeypatch.setattr( + helpers.asyncio, "get_running_loop", lambda: _Loop(fake_getaddrinfo) + ) + + target, host_header = await _validated_fetch_target("http://example.com/cat.png") + + assert target == "http://93.184.216.34/cat.png" + assert host_header == "example.com" + + +@pytest.mark.asyncio +async def test_https_target_keeps_hostname_for_tls( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def fake_getaddrinfo(*args: object, **kwargs: object) -> list[tuple]: + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443))] + + monkeypatch.setattr( + helpers.asyncio, "get_running_loop", lambda: _Loop(fake_getaddrinfo) + ) + + target, host_header = await _validated_fetch_target("https://example.com/cat.png") + + assert target == "https://example.com/cat.png" + assert host_header == "example.com" + + +@pytest.fixture +def reachable_sink(sink: str, monkeypatch: pytest.MonkeyPatch) -> str: + """Let the local sink stand in for a public host, so cap tests keep the + address validation intact instead of disabling it.""" + + async def passthrough(url: str) -> tuple[str, str]: + return url, "images.example.com" + + monkeypatch.setattr(helpers, "_validated_fetch_target", passthrough) + return sink + + +@pytest.mark.asyncio +async def test_downloaded_bytes_are_capped(reachable_sink: str) -> None: + body = await _fetch_image_from_url(f"{reachable_sink}/allowed") + + assert body is not None + assert len(body) <= IMAGE_FETCH_MAX_BYTES + assert REQUESTED_PATHS == ["/allowed"] + + +@pytest.mark.asyncio +async def test_url_fetches_are_capped_per_request(reachable_sink: str) -> None: + urls = IMAGE_FETCH_MAX_PER_REQUEST + 3 + messages = [ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": f"{reachable_sink}/{index}"}} + for index in range(urls) + ], + } + ] + + await estimate_image_tokens_in_messages(messages) + + assert len(REQUESTED_PATHS) == IMAGE_FETCH_MAX_PER_REQUEST diff --git a/tests/unit/test_lightning_settlement.py b/tests/unit/test_lightning_settlement.py index dbec3bd7..8d7c96a9 100644 --- a/tests/unit/test_lightning_settlement.py +++ b/tests/unit/test_lightning_settlement.py @@ -1,6 +1,6 @@ import asyncio import time -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Iterator from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, Mock, patch @@ -20,9 +20,17 @@ from routstr.lightning import ( get_invoice_status, recover_invoice, ) +from routstr.mint import MintRateGuard from routstr.wallet import Wallet +@pytest.fixture(autouse=True) +def _clear_mint_rate_guards() -> Iterator[None]: + MintRateGuard._guards.clear() + yield + MintRateGuard._guards.clear() + + def _invoice(**overrides: object) -> SimpleNamespace: values = { "id": "invoice-1", @@ -33,8 +41,6 @@ def _invoice(**overrides: object) -> SimpleNamespace: "paid_at": None, "api_key_hash": None, "mint_url": "http://mint:3338", - "balance_limit": None, - "balance_limit_reset": None, "validity_date": None, "created_at": 1, "expires_at": 2, diff --git a/tests/unit/test_lnurl_amount_and_destination.py b/tests/unit/test_lnurl_amount_and_destination.py new file mode 100644 index 00000000..72986500 --- /dev/null +++ b/tests/unit/test_lnurl_amount_and_destination.py @@ -0,0 +1,483 @@ +"""LNURL payments must verify the invoice amount and the destination. + +Findings 4 and 5 ship together: the amount check is what makes it safe to +hand an LNURL a set of unreserved proofs, and reserving only after that check +is what stops a pre-dispatch failure from stranding proofs. +""" + +import math +import socket +from collections.abc import Callable +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from cashu.core.base import MeltQuoteState + +from routstr.payment import lnurl as lnurl_module +from routstr.payment.lnurl import ( + LNURLError, + get_lnurl_data, + get_lnurl_invoice, + raw_send_to_lnurl, +) + +LNURL_DATA = { + "callback_url": "https://ln.tld/cb", + "min_sendable": 1_000, + "max_sendable": 100_000_000, +} + +EXPECTED_QUOTE_SAT = 999 + + +def _wallet( + quote_amount: int | None = None, +) -> tuple[MagicMock, list[MagicMock]]: + proofs = [MagicMock(amount=1000, reserved=False)] + wallet = MagicMock(url="https://mint.test") + wallet.get_fees_for_proofs.return_value = 0 + if quote_amount is None: + wallet.melt_quote = AsyncMock( + side_effect=[ + MagicMock(fee_reserve=1, quote="q", amount=1000), + MagicMock(fee_reserve=1, quote="q", amount=EXPECTED_QUOTE_SAT), + ] + ) + else: + wallet.melt_quote = AsyncMock( + return_value=MagicMock(fee_reserve=1, quote="q", amount=quote_amount) + ) + wallet.select_to_send = AsyncMock(return_value=(proofs, None)) + wallet.get_fees_for_proofs = MagicMock(return_value=0) + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid)) + wallet.set_reserved_for_send = AsyncMock() + return wallet, proofs + + +def _lnurl_patches() -> tuple[Any, Any]: + return ( + patch( + "routstr.payment.lnurl.get_lnurl_data", + AsyncMock(return_value=LNURL_DATA), + ), + patch( + "routstr.payment.lnurl.get_lnurl_invoice", + AsyncMock(return_value=("lnbc1...", {})), + ), + ) + + +def _mock_client(handler: Callable[[httpx.Request], httpx.Response]) -> Any: + real_client = httpx.AsyncClient + + def factory(*_args: object, **_kwargs: object) -> httpx.AsyncClient: + return real_client(transport=httpx.MockTransport(handler)) + + return patch.object(lnurl_module.httpx, "AsyncClient", factory) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "quote_amount", [EXPECTED_QUOTE_SAT + 1, EXPECTED_QUOTE_SAT * 5] +) +async def test_raw_send_to_lnurl_rejects_oversized_invoice(quote_amount: int) -> None: + wallet, proofs = _wallet(quote_amount) + data_patch, invoice_patch = _lnurl_patches() + + with data_patch, invoice_patch, pytest.raises(LNURLError, match="invoice amount"): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.select_to_send.assert_not_awaited() + wallet.melt.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_rejects_undersized_invoice() -> None: + wallet, proofs = _wallet(EXPECTED_QUOTE_SAT - 1) + data_patch, invoice_patch = _lnurl_patches() + + with data_patch, invoice_patch, pytest.raises(LNURLError, match="invoice amount"): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.select_to_send.assert_not_awaited() + wallet.melt.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_rejects_invoice_before_quote_checkpoint() -> None: + wallet, proofs = _wallet(EXPECTED_QUOTE_SAT * 2) + checkpoint = AsyncMock() + data_patch, invoice_patch = _lnurl_patches() + + with data_patch, invoice_patch, pytest.raises(LNURLError): + await raw_send_to_lnurl( + wallet, + proofs, + "owner@ln.tld", + "sat", + amount=1000, + on_melt_quote=checkpoint, + ) + + checkpoint.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_accepts_exact_invoice() -> None: + wallet, proofs = _wallet() + data_patch, invoice_patch = _lnurl_patches() + + with data_patch, invoice_patch: + paid = await raw_send_to_lnurl( + wallet, proofs, "owner@ln.tld", "sat", amount=1000 + ) + + assert paid == EXPECTED_QUOTE_SAT * 1000 + wallet.select_to_send.assert_not_awaited() + wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=True) + wallet.melt.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None: + wallet, proofs = _wallet() + proofs[0].amount = 1_000_000 + wallet.melt_quote = AsyncMock( + side_effect=[ + MagicMock(fee_reserve=1, quote="q", amount=1_000_000), + MagicMock(fee_reserve=1, quote="q", amount=999_999), + ] + ) + data_patch, invoice_patch = _lnurl_patches() + + with data_patch, invoice_patch: + paid = await raw_send_to_lnurl( + wallet, proofs, "owner@ln.tld", "msat", amount=1_000_000 + ) + + assert paid == 999_999 + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_requotes_for_exact_input_fees_without_recursion() -> ( + None +): + proofs = [MagicMock(amount=1, reserved=False) for _ in range(500)] + wallet = MagicMock(url="https://mint.test") + wallet.get_fees_for_proofs = MagicMock( + side_effect=lambda selected: math.ceil(len(selected) / 100) + ) + wallet.melt_quote = AsyncMock( + side_effect=[ + MagicMock(fee_reserve=10, quote="q1", amount=500), + MagicMock(fee_reserve=10, quote="q2", amount=485), + ] + ) + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid)) + wallet.set_reserved_for_send = AsyncMock() + checkpoint = AsyncMock() + data_patch, invoice_patch = _lnurl_patches() + + with data_patch, invoice_patch: + paid = await raw_send_to_lnurl( + wallet, + proofs, + "owner@ln.tld", + "sat", + amount=500, + on_melt_quote=checkpoint, + ) + + assert paid == 485_000 + assert wallet.melt_quote.await_count == 2 + checkpoint.assert_awaited_once_with("q2") + wallet.select_to_send.assert_not_called() + selected = wallet.melt.await_args.kwargs["proofs"] + assert sum(proof.amount for proof in selected) == 500 + assert 485 + 10 + wallet.get_fees_for_proofs(selected) == 500 + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_requires_an_explicit_amount() -> None: + wallet, proofs = _wallet() + data_patch, invoice_patch = _lnurl_patches() + + with data_patch, invoice_patch, pytest.raises(ValueError, match="amount"): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat") + + wallet.select_to_send.assert_not_awaited() + wallet.melt.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "address", + [ + "owner@127.0.0.1", + "owner@localhost", + "owner@10.0.0.5", + "owner@[::1]", + "http://ln.tld/lnurlp/owner", + ], +) +async def test_get_lnurl_data_rejects_non_public_destination(address: str) -> None: + requested: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requested.append(str(request.url)) + return httpx.Response( + 200, json={"tag": "payRequest", "callback": "https://x/y"} + ) + + with _mock_client(handler), pytest.raises(LNURLError): + await get_lnurl_data(address) + + assert requested == [] + + +@pytest.mark.asyncio +async def test_get_lnurl_data_rejects_redirect_to_private_host() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.host == "ln.tld": + return httpx.Response( + 302, headers={"location": "https://169.254.169.254/latest/meta-data"} + ) + return httpx.Response( + 200, json={"tag": "payRequest", "callback": "https://x/y"} + ) + + with _mock_client(handler), pytest.raises(LNURLError, match="destination"): + await get_lnurl_data("owner@ln.tld") + + +@pytest.mark.asyncio +async def test_get_lnurl_data_rejects_downgrade_redirect() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.scheme == "https": + return httpx.Response(302, headers={"location": "http://ln.tld/plain"}) + return httpx.Response( + 200, json={"tag": "payRequest", "callback": "https://x/y"} + ) + + with _mock_client(handler), pytest.raises(LNURLError, match="destination"): + await get_lnurl_data("owner@ln.tld") + + +@pytest.mark.asyncio +async def test_get_lnurl_data_rejects_private_callback_url() -> None: + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, json={"tag": "payRequest", "callback": "http://127.0.0.1:8000/cb"} + ) + + with _mock_client(handler), pytest.raises(LNURLError, match="destination"): + await get_lnurl_data("owner@ln.tld") + + +@pytest.mark.asyncio +async def test_get_lnurl_data_error_does_not_leak_response_body() -> None: + secret = "SUPERSECRETBODYMARKER" + + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"tag": secret, "callback": secret}) + + with _mock_client(handler), pytest.raises(LNURLError) as excinfo: + await get_lnurl_data("owner@ln.tld") + + assert secret not in str(excinfo.value) + + +@pytest.mark.asyncio +async def test_get_lnurl_invoice_error_does_not_leak_response_body() -> None: + secret = "SUPERSECRETBODYMARKER" + + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"reason": secret, "internal": secret}) + + with _mock_client(handler), pytest.raises(LNURLError) as excinfo: + await get_lnurl_invoice("https://ln.tld/cb", 1000) + + assert secret not in str(excinfo.value) + + +@pytest.mark.asyncio +async def test_get_lnurl_invoice_rejects_redirect_to_private_host() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.host == "ln.tld": + return httpx.Response(302, headers={"location": "https://192.168.1.1/cb"}) + return httpx.Response(200, json={"pr": "lnbc1..."}) + + with _mock_client(handler), pytest.raises(LNURLError, match="destination"): + await get_lnurl_invoice("https://ln.tld/cb", 1000) + + +@pytest.mark.asyncio +async def test_send_to_lnurl_does_not_reserve_before_lnurl_validation() -> None: + from routstr import wallet as wallet_module + + wallet, proofs = _wallet() + + with ( + patch.object( + wallet_module, + "find_trusted_mint_with_funds", + AsyncMock(return_value="https://mint.test"), + ), + patch.object(wallet_module, "get_wallet", AsyncMock(return_value=wallet)), + patch.object( + wallet_module, + "get_proofs_per_mint_and_unit", + MagicMock(return_value=proofs), + ), + patch.object( + wallet_module, + "raw_send_to_lnurl", + AsyncMock(side_effect=LNURLError("destination rejected")), + ) as raw_send, + pytest.raises(LNURLError), + ): + await wallet_module.send_to_lnurl( + 1000, "sat", "https://mint.test", "owner@ln.tld" + ) + + wallet.select_to_send.assert_not_awaited() + assert raw_send.await_args is not None + assert raw_send.await_args.args[1] is proofs + assert raw_send.await_args.kwargs["amount"] == 1000 + + +def _patch_getaddrinfo(ip: str) -> Any: + """Force DNS resolution of any hostname to a single fixed IP.""" + + async def fake_getaddrinfo(host: str, port: int, **_kw: object) -> list[Any]: + return [ + (socket.AF_INET, socket.SOCK_STREAM, socket.IPPROTO_TCP, "", (ip, port)) + ] + + loop = MagicMock() + loop.getaddrinfo = fake_getaddrinfo + return patch.object( + lnurl_module.asyncio, "get_running_loop", return_value=loop + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "private_ip", ["127.0.0.1", "10.0.0.5", "169.254.169.254", "192.168.1.1"] +) +async def test_guard_rejects_public_hostname_resolving_to_private( + private_ip: str, +) -> None: + """A public-looking name must be rejected when DNS points it inward (SSRF).""" + with ( + _patch_getaddrinfo(private_ip), + pytest.raises(LNURLError, match="public host"), + ): + await lnurl_module._require_public_https_destination( + httpx.URL("https://totally-public.example.com/cb") + ) + + +@pytest.mark.asyncio +async def test_guard_allows_public_hostname_resolving_to_public() -> None: + with _patch_getaddrinfo("93.184.216.34"): + await lnurl_module._require_public_https_destination( + httpx.URL("https://totally-public.example.com/cb") + ) + + +@pytest.mark.asyncio +async def test_get_lnurl_data_rejects_oversized_response() -> None: + big = b"x" * (lnurl_module._MAX_LNURL_RESPONSE_BYTES + 1) + + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response(200, content=big) + + with ( + _patch_getaddrinfo("93.184.216.34"), + _mock_client(handler), + pytest.raises(LNURLError, match="size limit"), + ): + await get_lnurl_data("owner@ln.tld") + + +def test_select_melt_proofs_ignores_fees_for_unneeded_wallet_proofs() -> None: + from routstr.payment.lnurl import _select_melt_proofs + + wallet = MagicMock() + wallet.get_fees_for_proofs = MagicMock(side_effect=lambda selected: len(selected)) + proofs = [MagicMock(amount=2048, reserved=False) for _ in range(1100)] + + selected, shortfall = _select_melt_proofs( + wallet, + proofs, + quote_amount=1061, + fee_reserve=1, + gross_budget=1061, + ) + + assert selected is None + assert shortfall == 2 + assert wallet.get_fees_for_proofs.call_count == 1 + + +def test_select_melt_proofs_respects_mint_input_limit() -> None: + from cashu.core.settings import settings as cashu_settings + + from routstr.payment.lnurl import _select_melt_proofs + + limit = cashu_settings.mint_max_request_length + wallet = MagicMock() + wallet.get_fees_for_proofs = MagicMock(return_value=0) + proofs = [MagicMock(amount=1, reserved=False) for _ in range(limit + 563)] + + selected, shortfall = _select_melt_proofs( + wallet, + proofs, + quote_amount=limit + 563, + fee_reserve=0, + gross_budget=limit + 563, + ) + + assert selected is None + assert shortfall == 563 + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_pays_what_the_input_limit_allows() -> None: + from cashu.core.settings import settings as cashu_settings + + limit = cashu_settings.mint_max_request_length + proofs = [MagicMock(amount=1, reserved=False) for _ in range(limit + 563)] + wallet = MagicMock(url="https://mint.test") + wallet.get_fees_for_proofs = MagicMock(return_value=0) + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid)) + wallet.set_reserved_for_send = AsyncMock() + + requested: list[int] = [] + + async def invoice(_callback: str, amount_msat: int) -> tuple[str, dict]: + requested.append(amount_msat) + return "lnbc1...", {} + + async def melt_quote(invoice: str) -> MagicMock: + return MagicMock(fee_reserve=0, quote="q", amount=requested[-1] // 1000) + + wallet.melt_quote = AsyncMock(side_effect=melt_quote) + + with ( + patch( + "routstr.payment.lnurl.get_lnurl_data", AsyncMock(return_value=LNURL_DATA) + ), + patch( + "routstr.payment.lnurl.get_lnurl_invoice", AsyncMock(side_effect=invoice) + ), + ): + paid = await raw_send_to_lnurl( + wallet, proofs, "owner@ln.tld", "sat", amount=limit + 563 + ) + + assert paid == limit * 1000 + assert len(wallet.melt.await_args.kwargs["proofs"]) == limit diff --git a/tests/unit/test_lnurl_melt_timeout.py b/tests/unit/test_lnurl_melt_timeout.py index a370a7a8..ede3f5af 100644 --- a/tests/unit/test_lnurl_melt_timeout.py +++ b/tests/unit/test_lnurl_melt_timeout.py @@ -1,6 +1,7 @@ """LNURL melt attempts must not misclassify ambiguous payment outcomes.""" import asyncio +from collections.abc import Iterator from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -11,10 +12,19 @@ from cashu.core.base import MeltQuoteState from routstr.core.settings import settings from routstr.mint import MintCooldownError, MintRateGuard from routstr.payment.lnurl import ( + LNURLError, MeltOutcomeAmbiguousError, raw_send_to_lnurl, ) + +@pytest.fixture(autouse=True) +def _clear_mint_guards() -> Iterator[None]: + MintRateGuard._guards.clear() + yield + MintRateGuard._guards.clear() + + LNURL_DATA = { "callback_url": "https://ln.tld/cb", "min_sendable": 1_000, @@ -22,11 +32,23 @@ LNURL_DATA = { } +QUOTE_AMOUNT_SAT = 999 + + def _wallet() -> tuple[MagicMock, list[MagicMock]]: - proofs = [MagicMock(amount=1000)] + proofs = [MagicMock(amount=1000, reserved=False)] wallet = MagicMock(url="https://mint.test") - wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q")) + wallet.melt_quote = AsyncMock( + side_effect=[ + MagicMock(fee_reserve=1, quote="q", amount=1000), + MagicMock(fee_reserve=1, quote="q", amount=QUOTE_AMOUNT_SAT), + ] + ) + wallet.melt = AsyncMock() wallet.select_to_send = AsyncMock(return_value=(proofs, None)) + wallet.get_fees_for_proofs = MagicMock(return_value=0) + wallet.set_reserved_for_melt = AsyncMock() + wallet.set_reserved_for_send = AsyncMock() return wallet, proofs @@ -44,7 +66,26 @@ def _lnurl_patches() -> tuple[Any, Any]: @pytest.mark.asyncio -async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> None: +async def test_raw_send_to_lnurl_direct_unpaid_is_retry_safe() -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.unpaid)) + wallet.get_melt_quote = AsyncMock() + data_patch, invoice_patch = _lnurl_patches() + + with ( + data_patch, + invoice_patch, + pytest.raises(LNURLError, match="confirmed that the melt was unpaid") as raised, + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + assert not isinstance(raised.value, MeltOutcomeAmbiguousError) + wallet.get_melt_quote.assert_not_awaited() + wallet.set_reserved_for_send.assert_any_await(proofs, reserved=False) + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_timeout_then_unpaid_remains_ambiguous() -> None: wallet, proofs = _wallet() async def _hang(**kwargs: object) -> None: @@ -61,12 +102,65 @@ async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> Non patch.object(settings, "mint_retry_max_attempts", 0), data_patch, invoice_patch, - pytest.raises(MeltOutcomeAmbiguousError, match="outcome is ambiguous"), + pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"), ): await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) wallet.get_melt_quote.assert_awaited_once_with("q") - wallet.set_reserved_for_melt.assert_not_called() + assert wallet.set_reserved_for_melt.await_count == 2 + wallet.set_reserved_for_melt.assert_awaited_with( + proofs, reserved=True, quote_id="q" + ) + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_wrapped_transport_unpaid_remains_ambiguous() -> None: + wallet, proofs = _wallet() + + async def _wrapped_transport_error(**kwargs: object) -> None: + try: + raise httpx.ReadTimeout("response lost") + except httpx.ReadTimeout as transport_error: + raise Exception("could not pay invoice") from transport_error + + wallet.melt = AsyncMock(side_effect=_wrapped_transport_error) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.unpaid) + ) + data_patch, invoice_patch = _lnurl_patches() + + with ( + patch.object(settings, "mint_retry_max_attempts", 3), + data_patch, + invoice_patch, + pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.melt.assert_awaited_once() + wallet.get_melt_quote.assert_awaited_once_with("q") + assert wallet.set_reserved_for_melt.await_count == 2 + wallet.set_reserved_for_melt.assert_awaited_with( + proofs, reserved=True, quote_id="q" + ) + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_does_not_retry_melt_quote_timeout() -> None: + wallet, proofs = _wallet() + wallet.melt_quote = AsyncMock(side_effect=httpx.ReadTimeout("response lost")) + data_patch, invoice_patch = _lnurl_patches() + + with ( + patch.object(settings, "mint_retry_max_attempts", 3), + data_patch, + invoice_patch, + pytest.raises(httpx.TimeoutException), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.melt_quote.assert_awaited_once() + wallet.melt.assert_not_awaited() @pytest.mark.asyncio @@ -92,6 +186,9 @@ async def test_raw_send_to_lnurl_timeout_reconciled_paid_is_success() -> None: assert paid > 0 wallet.get_melt_quote.assert_awaited_once_with("q") + wallet.set_reserved_for_melt.assert_awaited_once_with( + proofs, reserved=True, quote_id="q" + ) @pytest.mark.asyncio @@ -114,6 +211,45 @@ async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None: wallet.get_melt_quote.assert_awaited_once_with("q") +@pytest.mark.asyncio +async def test_pending_then_immediate_unpaid_remains_reserved_and_ambiguous() -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending)) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.unpaid) + ) + data_patch, invoice_patch = _lnurl_patches() + + with ( + data_patch, + invoice_patch, + pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.set_reserved_for_melt.assert_awaited_once_with( + proofs, reserved=True, quote_id="q" + ) + + +@pytest.mark.asyncio +async def test_immediate_unpaid_reservation_failure_stays_ambiguous() -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending)) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.unpaid) + ) + wallet.set_reserved_for_melt = AsyncMock(side_effect=OSError("db locked")) + data_patch, invoice_patch = _lnurl_patches() + + with ( + data_patch, + invoice_patch, + pytest.raises(MeltOutcomeAmbiguousError, match="could not be restored"), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + @pytest.mark.asyncio @pytest.mark.parametrize("rate_error", ["cooldown", "http_429"]) async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs( @@ -147,7 +283,8 @@ async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs( await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) wallet.melt.assert_not_awaited() - wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=False) + assert wallet.set_reserved_for_send.await_count == 2 + wallet.set_reserved_for_send.assert_awaited_with(proofs, reserved=False) @pytest.mark.asyncio @@ -172,7 +309,8 @@ async def test_real_mint_wrapper_http_429_unreserves_proofs() -> None: await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) wallet.melt.assert_awaited_once() - wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=False) + assert wallet.set_reserved_for_send.await_count == 2 + wallet.set_reserved_for_send.assert_awaited_with(proofs, reserved=False) MintRateGuard._guards.pop(str(wallet.url), None) diff --git a/tests/unit/test_log_secret_redaction.py b/tests/unit/test_log_secret_redaction.py new file mode 100644 index 00000000..0f660dda --- /dev/null +++ b/tests/unit/test_log_secret_redaction.py @@ -0,0 +1,238 @@ +"""Regression tests for spendable credentials leaking into the dated log files. + +Everything here asserts against the bytes the ``DailyRotatingFileHandler`` +actually wrote to disk. Asserting against a mock would pass even if the JSON +formatter emitted the raw ``extra`` dict, which is exactly the bug. +""" + +import json +import logging +import os +import time +from collections.abc import Callable, Iterator +from pathlib import Path +from typing import Any + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from pythonjsonlogger import jsonlogger + +from routstr.core.logging import ( + DailyRotatingFileHandler, + RequestIdFilter, + SecurityFilter, + VersionFilter, +) +from routstr.core.middleware import LoggingMiddleware + +REFUND_TOKEN = ( + "cashuBo2FteCJodHRwczovL21pbnQubWluaWJpdHMuY2FzaC9CaXRjb2luYXVjc" + "2F0YXSBomFpSAA5tMOFA4EXYXCBo2FhAmFzeEA5NmY0NTFhZjMzMGY3ZmM2ZGY5" +) +HASHED_KEY = "b3d9f1c2a8574e60b7c1f0aa9d2e4c85f6b70a1932de84cc57bf90ae1d2c3f47" + + +@pytest.fixture +def log_dir(tmp_path: Path) -> Path: + directory = tmp_path / "logs" + directory.mkdir() + return directory + + +@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", + interval=1, + backupCount=30, + ) + handler.setLevel(logging.DEBUG) + handler.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 log_filter in (VersionFilter(), RequestIdFilter(), SecurityFilter()): + handler.addFilter(log_filter) + try: + yield handler + finally: + handler.close() + + +@pytest.fixture +def emit(handler: DailyRotatingFileHandler) -> Callable[..., str]: + """Log one record and return the raw text of the dated file it landed in.""" + logger = logging.getLogger("routstr.test.redaction") + logger.setLevel(logging.DEBUG) + logger.propagate = False + logger.handlers = [handler] + + def _emit(message: str, **extra: Any) -> str: + logger.info(message, extra=extra) + handler.flush() + return Path(handler.baseFilename).read_text() + + return _emit + + +def test_refund_token_never_reaches_the_dated_file(emit: Callable[..., str]) -> None: + written = emit( + "refund_wallet_endpoint: cashu token issued", + token=REFUND_TOKEN, + amount=1500, + currency="sat", + ) + assert REFUND_TOKEN not in written + assert "1500" in written + + +def test_authorization_values_never_reach_the_dated_file( + emit: Callable[..., str], +) -> None: + written = emit( + "Incoming request", + authorization=f"Bearer sk-{HASHED_KEY}", + headers={"Authorization": f"Bearer sk-{HASHED_KEY}"}, + ) + assert HASHED_KEY not in written + assert "Bearer sk-" not in written + + +def test_full_key_hashes_never_reach_the_dated_file(emit: Callable[..., str]) -> None: + written = emit( + "refund_wallet_endpoint: balance restored after mint failure", + hashed_key=HASHED_KEY, + key_hash=HASHED_KEY, + restored_balance=42, + ) + assert HASHED_KEY not in written + assert "42" in written + + +def test_secrets_nested_in_dicts_and_lists_never_reach_the_dated_file( + emit: Callable[..., str], +) -> None: + written = emit( + "Upstream call failed", + context={ + "attempts": [ + {"headers": {"authorization": f"Bearer sk-{HASHED_KEY}"}}, + {"body": {"refund": {"token": REFUND_TOKEN}}}, + ], + "provider": "openai", + }, + ) + assert HASHED_KEY not in written + assert REFUND_TOKEN not in written + assert "openai" in written + + +def test_query_string_secrets_never_reach_the_dated_file( + emit: Callable[..., str], +) -> None: + written = emit( + "Incoming request", + path="/v1/wallet/refund", + query_params={"api_key": f"sk-{HASHED_KEY}", "page": "2"}, + target=f"/v1/wallet/refund?token={REFUND_TOKEN}", + ) + assert HASHED_KEY not in written + assert REFUND_TOKEN not in written + assert "/v1/wallet/refund" in written + + +def test_benign_telemetry_is_not_redacted(emit: Callable[..., str]) -> None: + written = emit( + "Request completed", + method="POST", + path="/v1/chat/completions", + model="gpt-4o-mini", + status_code=200, + duration_ms=13.5, + input_tokens=120, + key_hash=HASHED_KEY[:8], + mint_url="https://mint.minibits.cash/Bitcoin", + ) + record = json.loads(written.strip().splitlines()[-1]) + assert record["model"] == "gpt-4o-mini" + assert record["status_code"] == 200 + assert record["duration_ms"] == 13.5 + assert record["input_tokens"] == 120 + assert record["key_hash"] == HASHED_KEY[:8] + assert record["mint_url"] == "https://mint.minibits.cash/Bitcoin" + assert record["message"] == "Request completed" + + +def test_self_referential_extra_does_not_hang_the_logger( + emit: Callable[..., str], +) -> None: + cyclic: dict[str, Any] = {"token": REFUND_TOKEN} + cyclic["self"] = cyclic + deep: dict[str, Any] = {"token": REFUND_TOKEN} + for _ in range(200): + deep = {"nested": deep} + + written = emit("Cyclic payload", context=cyclic, deep=deep) + + assert REFUND_TOKEN not in written + assert json.loads(written.strip().splitlines()[-1])["message"] == "Cyclic payload" + + +def test_middleware_logs_query_param_names_without_values( + handler: DailyRotatingFileHandler, +) -> None: + app = FastAPI() + app.add_middleware(LoggingMiddleware) + + @app.get("/v1/wallet/refund") + async def refund() -> dict[str, str]: + return {"status": "ok"} + + middleware_logger = logging.getLogger("routstr.core.middleware") + middleware_logger.setLevel(logging.INFO) + middleware_logger.propagate = False + original_handlers = middleware_logger.handlers + middleware_logger.handlers = [handler] + try: + with TestClient(app) as client: + response = client.get( + "/v1/wallet/refund", params={"api_key": f"sk-{HASHED_KEY}", "page": "2"} + ) + assert response.status_code == 200 + finally: + middleware_logger.handlers = original_handlers + + handler.flush() + written = Path(handler.baseFilename).read_text() + assert HASHED_KEY not in written + record = json.loads(written.strip().splitlines()[0]) + assert record["path"] == "/v1/wallet/refund" + assert record["query_param_names"] == ["api_key", "page"] + + +def test_forced_rollover_enforces_the_retention_limit( + handler: DailyRotatingFileHandler, log_dir: Path +) -> None: + handler.backupCount = 3 + handler.emit( + logging.LogRecord("t", logging.INFO, "", 0, "current", (), None), + ) + handler.flush() + + now = time.time() + for day in range(1, 6): + stale = log_dir / f"app_2024-01-0{day}.log" + stale.write_text("stale\n") + os.utime(stale, (now - day * 86400, now - day * 86400)) + + handler.doRollover() + + remaining = sorted(p.name for p in log_dir.glob("app_*.log")) + assert len(remaining) == handler.backupCount + assert Path(handler.baseFilename).name in remaining diff --git a/tests/unit/test_melt_reconciliation.py b/tests/unit/test_melt_reconciliation.py deleted file mode 100644 index a68cb64b..00000000 --- a/tests/unit/test_melt_reconciliation.py +++ /dev/null @@ -1,91 +0,0 @@ -from unittest.mock import AsyncMock, Mock - -import pytest -from cashu.core.base import MeltQuoteState, ProofSpentState - -from routstr.wallet import ( - TokenConsumedError, - _confirm_melt_paid, - _reconcile_ambiguous_melt, -) - - -@pytest.mark.asyncio -async def test_paid_quote_is_authoritative_when_proof_lookup_would_fail() -> None: - wallet = Mock( - url="http://source-mint:3338", - get_melt_quote=AsyncMock(return_value=Mock(state=MeltQuoteState.paid)), - check_proof_state=AsyncMock(side_effect=RuntimeError("proof API unavailable")), - ) - - assert await _reconcile_ambiguous_melt(wallet, "quote-1", [Mock()]) is True - wallet.check_proof_state.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_timeout_snapshot_unpaid_unspent_remains_non_retryable() -> None: - wallet = Mock( - url="http://source-mint:3338", - get_melt_quote=AsyncMock(return_value=Mock(state=MeltQuoteState.unpaid)), - check_proof_state=AsyncMock( - return_value=Mock(states=[Mock(state=ProofSpentState.unspent)]) - ), - ) - - with pytest.raises(TokenConsumedError, match="ambiguous"): - await _reconcile_ambiguous_melt(wallet, "quote-2", [Mock()]) - - -@pytest.mark.asyncio -async def test_successful_pending_melt_response_requires_reconciliation() -> None: - wallet = Mock( - url="http://source-mint:3338", - get_melt_quote=AsyncMock(return_value=Mock(state=MeltQuoteState.pending)), - check_proof_state=AsyncMock( - return_value=Mock(states=[Mock(state=ProofSpentState.pending)]) - ), - ) - - with pytest.raises(TokenConsumedError, match="ambiguous"): - await _confirm_melt_paid( - wallet, - "quote-pending", - [Mock()], - Mock(state=MeltQuoteState.pending), - ) - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("quote_state", "proof_state"), - [ - (MeltQuoteState.pending, ProofSpentState.pending), - (MeltQuoteState.unpaid, ProofSpentState.spent), - (MeltQuoteState.unpaid, ProofSpentState.pending), - ], -) -async def test_ambiguous_or_consumed_melt_is_never_reported_unspent( - quote_state: MeltQuoteState, proof_state: ProofSpentState -) -> None: - wallet = Mock( - url="http://source-mint:3338", - get_melt_quote=AsyncMock(return_value=Mock(state=quote_state)), - check_proof_state=AsyncMock( - return_value=Mock(states=[Mock(state=proof_state)]) - ), - ) - - with pytest.raises(TokenConsumedError, match="reconciliation required"): - await _reconcile_ambiguous_melt(wallet, "quote-3", [Mock()]) - - -@pytest.mark.asyncio -async def test_failed_melt_reconciliation_is_non_retryable() -> None: - wallet = Mock( - url="http://source-mint:3338", - get_melt_quote=AsyncMock(side_effect=RuntimeError("mint unavailable")), - check_proof_state=AsyncMock(), - ) - - with pytest.raises(TokenConsumedError, match="outcome is unknown"): - await _reconcile_ambiguous_melt(wallet, "quote-4", [Mock()]) diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index cb3fe9d7..294f5c0c 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -305,6 +305,10 @@ async def test_dispatch_strips_anthropic_only_fields_before_litellm() -> None: "service_tier": "auto", "anthropic_beta": "abc", "anthropic_version": "2023-06-01", + "api_base": "https://attacker.invalid", + "api_key": "client-controlled-key", + "custom_llm_provider": "client-controlled-provider", + "unexpected_field": "must-not-leak", } ).encode() @@ -342,6 +346,8 @@ async def test_dispatch_strips_anthropic_only_fields_before_litellm() -> None: "service_tier", "anthropic_beta", "anthropic_version", + "custom_llm_provider", + "unexpected_field", ): assert stripped not in forwarded, ( f"Anthropic-only field {stripped!r} leaked through to litellm" @@ -349,6 +355,64 @@ async def test_dispatch_strips_anthropic_only_fields_before_litellm() -> None: # Core fields preserved assert forwarded["max_tokens"] == 64 assert forwarded["messages"] == [{"role": "user", "content": "hi"}] + # Dispatch-controlled values cannot be overridden by the client body. + assert forwarded["model"] == "openai/openai/gpt-4o-mini" + assert forwarded["api_base"] == "http://test" + assert forwarded["api_key"] == "upstream-key" + assert forwarded["stream"] is True + + +@pytest.mark.asyncio +async def test_dispatch_forwards_all_allowlisted_messages_fields() -> None: + provider = _make_provider() + key = _make_key() + model = _make_model() + session = _make_session() + allowed = { + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 64, + "system": "Be concise", + "temperature": 0.2, + "top_p": 0.9, + "top_k": 20, + "stop_sequences": ["STOP"], + "tools": [ + { + "name": "lookup", + "description": "Look something up", + "input_schema": {"type": "object", "properties": {}}, + } + ], + "tool_choice": {"type": "tool", "name": "lookup"}, + "metadata": {"user_id": "test-user"}, + } + body = json.dumps({"model": model.id, **allowed}).encode() + captured: dict[str, Any] = {} + + async def fake_acreate(**kwargs: Any) -> dict: + captured["kwargs"] = kwargs + return _anthropic_non_stream_response() + + with ( + patch( + "litellm.anthropic.messages.acreate", + new=AsyncMock(side_effect=fake_acreate), + ), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + new=AsyncMock(return_value={"total_msats": 0, "total_usd": 0.0}), + ), + ): + await provider._forward_messages_via_litellm( + request_body=body, + key=key, + session=session, + max_cost_for_model=10_000, + model_obj=model, + ) + + forwarded = captured["kwargs"] + assert {field: forwarded[field] for field in allowed} == allowed @pytest.mark.asyncio @@ -1070,18 +1134,30 @@ async def test_forward_x_cashu_request_handles_count_tokens_locally() -> None: "prepare_request_body", side_effect=AssertionError("upstream should not be called"), ): - response = await provider.forward_x_cashu_request( - request=request, - path="v1/messages/count_tokens", - headers={}, - amount=5_000, - unit="sat", - max_cost_for_model=10_000, - model_obj=model, - mint="https://mint", - ) + with patch.object( + provider, + "send_refund", + new=AsyncMock(return_value="refund-token"), + ) as send_refund: + response = await provider.forward_x_cashu_request( + request=request, + path="v1/messages/count_tokens", + headers={}, + amount=5_000, + unit="sat", + max_cost_for_model=10_000, + model_obj=model, + mint="https://mint", + ) + send_refund.assert_awaited_once_with( + 5_000, + "sat", + "https://mint", + request_id="req-test", + ) assert response.status_code == 200 + assert response.headers["X-Cashu"] == "refund-token" body = response.body if isinstance(response.body, bytes) else bytes(response.body) payload = json.loads(body.decode()) assert "input_tokens" in payload @@ -1361,15 +1437,34 @@ async def test_dispatch_uses_url_detected_prefix_for_fireworks_custom_row() -> N ["handle_x_cashu", "handle_x_cashu_responses"], ) @pytest.mark.parametrize( - "error", + ("error", "expected_type", "expected_code", "expected_message"), [ - httpx.ConnectError("All connection attempts failed"), - MintConnectionError("Cashu mint is unreachable"), - TimeoutError("timed out connecting to mint"), + ( + httpx.ConnectError("All connection attempts failed"), + "mint_unreachable", + "cashu_mint_unreachable", + "Cashu mint is unreachable; retry later", + ), + ( + MintConnectionError("connect to http://mint:3338 refused"), + "mint_unreachable", + "cashu_mint_unreachable", + "Cashu mint is unreachable; retry later", + ), + ( + TimeoutError("timed out connecting to mint"), + "mint_timeout", + "cashu_mint_timeout", + "Cashu mint did not respond in time; retry later", + ), ], ) async def test_x_cashu_mint_unreachable_returns_503( - handler_name: str, error: Exception + handler_name: str, + error: Exception, + expected_type: str, + expected_code: str, + expected_message: str, ) -> None: """Both X-Cashu entrypoints classify a down mint as 503 mint_unreachable, not a generic 400 cashu_error.""" @@ -1391,9 +1486,9 @@ async def test_x_cashu_mint_unreachable_returns_503( assert response.status_code == 503 body = json.loads(bytes(response.body)) - assert body["error"]["type"] == "mint_unreachable" - assert body["error"]["message"] == "Cashu mint is unreachable" - assert body["error"]["code"] == "cashu_mint_unreachable" + assert body["error"]["type"] == expected_type + assert body["error"]["message"] == expected_message + assert body["error"]["code"] == expected_code if str(error) != body["error"]["message"]: assert str(error) not in body["error"]["message"] @@ -1437,13 +1532,6 @@ async def test_x_cashu_mint_unreachable_returns_503( "Token value is too small to cover swap fees", "cashu_token_swap_fees_exceed_amount", ), - ( - ValueError("Failed to melt token from foreign mint http://m: boom"), - 422, - "mint_error", - "Failed to swap token from foreign mint", - "cashu_foreign_mint_swap_failed", - ), ( ValueError("some unexpected wallet condition"), 400, diff --git a/tests/unit/test_mint.py b/tests/unit/test_mint.py index 6b70a558..f6873fa4 100644 --- a/tests/unit/test_mint.py +++ b/tests/unit/test_mint.py @@ -10,6 +10,7 @@ from routstr.mint import ( MintRateGuard, MintRateLimitedError, fail_fast_mint_operations, + run_mint_operation, ) from routstr.wallet import Wallet @@ -103,6 +104,82 @@ async def test_cashu_429_dispatches_through_wallet_override() -> None: await wallet.mint_quote(1, Unit.sat) +@pytest.mark.asyncio +async def test_wrapped_transport_failure_opens_central_cooldown() -> None: + mint_url = "https://transport-failure.test" + MintRateGuard._guards.pop(mint_url, None) + + async def wrapped_failure() -> None: + try: + raise httpx.ReadTimeout("body stalled") + except httpx.ReadTimeout as error: + raise Exception("wallet wrapper") from error + + with pytest.raises(Exception, match="wallet wrapper"): + await run_mint_operation( + wrapped_failure, + mint_url=mint_url, + retry_timeouts=False, + ) + + guard = MintRateGuard.get(mint_url) + assert guard.cooldown_remaining() > 29 + probe = AsyncMock() + async with fail_fast_mint_operations(): + with pytest.raises(MintCooldownError): + await guard.run(probe) + probe.assert_not_awaited() + MintRateGuard._guards.pop(mint_url, None) + + +@pytest.mark.asyncio +async def test_timeout_retry_succeeds_without_opening_cooldown() -> None: + from routstr.core.settings import settings + + mint_url = "https://retryable-timeout.test" + MintRateGuard._guards.pop(mint_url, None) + calls = 0 + + async def flaky() -> str: + nonlocal calls + calls += 1 + if calls == 1: + raise httpx.ReadTimeout("first attempt stalled") + return "ok" + + with ( + patch.object(settings, "mint_retry_max_attempts", 2), + patch("routstr.mint.asyncio.sleep", AsyncMock()), + ): + result = await run_mint_operation(flaky, mint_url=mint_url) + + assert result == "ok" + assert calls == 2 + assert MintRateGuard.get(mint_url).cooldown_remaining() == 0.0 + MintRateGuard._guards.pop(mint_url, None) + + +@pytest.mark.asyncio +async def test_exhausted_timeout_retries_open_transport_cooldown() -> None: + from routstr.core.settings import settings + + mint_url = "https://exhausted-timeout.test" + MintRateGuard._guards.pop(mint_url, None) + + async def always_timeout() -> None: + raise httpx.ReadTimeout("stalled") + + with ( + patch.object(settings, "mint_retry_max_attempts", 1), + patch("routstr.mint.asyncio.sleep", AsyncMock()), + pytest.raises(httpx.TimeoutException), + ): + await run_mint_operation(always_timeout, mint_url=mint_url) + + assert MintRateGuard.get(mint_url).cooldown_remaining() > 29 + MintRateGuard._guards.pop(mint_url, None) + + async def test_guard_concurrency_change_preserves_cooldown_state() -> None: from routstr.core.settings import settings diff --git a/tests/unit/test_mint_fallback_trust.py b/tests/unit/test_mint_fallback_trust.py index b72314fc..6eec7111 100644 --- a/tests/unit/test_mint_fallback_trust.py +++ b/tests/unit/test_mint_fallback_trust.py @@ -1,7 +1,8 @@ """Persisted mint preferences must not bypass the configured trusted set.""" -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from routstr.core.settings import settings @@ -32,6 +33,23 @@ async def test_untrusted_allowed_mints_fall_back_to_trusted_set() -> None: assert attempted == [TRUSTED] +async def test_mint_quote_timeout_is_not_retried() -> None: + wallet = MagicMock() + wallet.request_mint = AsyncMock(side_effect=httpx.ReadTimeout("response lost")) + + with ( + patch.object(settings, "primary_mint", TRUSTED), + patch.object(settings, "cashu_mints", [TRUSTED]), + patch.object(settings, "mint_retry_max_attempts", 3), + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.lightning.mint_cooldown_remaining", return_value=0.0), + pytest.raises(Exception), + ): + await _request_mint_with_fallback(10) + + wallet.request_mint.assert_awaited_once_with(10) + + async def test_trusted_allowed_mints_are_used_verbatim() -> None: attempted: list[str] = [] diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py new file mode 100644 index 00000000..0b713652 --- /dev/null +++ b/tests/unit/test_model_path_routing.py @@ -0,0 +1,625 @@ +"""Model-path routing and fail-closed behavior.""" + +import json +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from routstr import proxy as proxy_module +from routstr.auth import ReservationSnapshot +from routstr.core.db import ApiKey +from routstr.upstream.model_paths import decode_model_path, encode_model_path + +MODEL_ID = "test-model" + + +def _make_upstream(db_id: int, status_code: int = 200) -> MagicMock: + upstream = MagicMock() + upstream.db_id = db_id + upstream.base_url = "http://localhost" + upstream.provider_type = f"provider-{db_id}" + upstream.supports_ehbp = False + upstream.prepare_headers = MagicMock(side_effect=lambda h: h) + upstream.on_upstream_error_redirect = AsyncMock() + upstream.forward_request = AsyncMock( + return_value=MagicMock(status_code=status_code, body=b"{}") + ) + return upstream + + +def _make_request(headers: dict[str, str], body: bytes) -> MagicMock: + request = MagicMock() + request.method = "POST" + request.headers = headers + request.body = AsyncMock(return_value=body) + request.state = MagicMock() + request.state.request_id = "req-model-path" + return request + + +async def _run_proxy( + request: MagicMock, + candidates: list[tuple[Any, Any]], + path: str = "v1/chat/completions", +) -> Any: + key = ApiKey(hashed_key="mpkey", balance=10_000) + reservation = ReservationSnapshot( + release_id="model-path-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=1_000)), + patch.object( + proxy_module, + "get_reservation_snapshot", + AsyncMock(return_value=reservation), + ), + patch.object(proxy_module, "revert_pay_for_request", AsyncMock()), + ): + return await proxy_module.proxy(request, path, session=MagicMock()) + + +def test_decode_model_path_round_trips_encode() -> None: + selector = decode_model_path( + encode_model_path("https://openrouter.ai/api/v1", 7, MODEL_ID, "deepinfra/fp8") + ) + assert selector is not None + assert selector.base_url == "https://openrouter.ai/api/v1" + assert selector.provider_id == 7 + assert selector.model_id == MODEL_ID + assert selector.endpoint_tag == "deepinfra/fp8" + + +def test_decode_model_path_without_endpoint_has_no_tag() -> None: + selector = decode_model_path(encode_model_path("http://localhost", 1, MODEL_ID)) + assert selector is not None + assert selector.endpoint_tag is None + + +@pytest.mark.parametrize( + "path", + [ + "", + "url=http://localhost&model-id=test-model", + "url=http://localhost&provider-id=abc&model-id=test-model", + "provider-id=1&model-id=test-model", + "url=http://localhost&provider-id=1", + ], +) +def test_decode_model_path_rejects_malformed_selectors(path: str) -> None: + assert decode_model_path(path) is None + + +@pytest.mark.asyncio +async def test_model_path_routes_to_the_selected_provider() -> None: + first, selected = _make_upstream(1), _make_upstream(2) + request = _make_request( + { + "authorization": "Bearer sk-mpkey", + "x-routstr-model-path": encode_model_path("http://localhost", 2, MODEL_ID), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + + await _run_proxy(request, [(MagicMock(), first), (MagicMock(), selected)]) + + selected.forward_request.assert_awaited_once() + first.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status_code", [400, 429, 500, 502, 503]) +async def test_model_path_failure_is_returned_without_falling_back( + status_code: int, +) -> None: + selected = _make_upstream(1, status_code=status_code) + fallback = _make_upstream(2) + request = _make_request( + { + "authorization": "Bearer sk-mpkey", + "x-routstr-model-path": encode_model_path("http://localhost", 1, MODEL_ID), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + + response = await _run_proxy( + request, [(MagicMock(), selected), (MagicMock(), fallback)] + ) + + assert response.status_code == status_code + selected.forward_request.assert_awaited_once() + fallback.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_unknown_provider_in_model_path_is_rejected() -> None: + upstream = _make_upstream(1) + request = _make_request( + { + "authorization": "Bearer sk-mpkey", + "x-routstr-model-path": encode_model_path("http://localhost", 99, MODEL_ID), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + + response = await _run_proxy(request, [(MagicMock(), upstream)]) + + assert response.status_code == 404 + assert json.loads(bytes(response.body))["error"]["type"] == "invalid_model_path" + upstream.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_model_path_disagreeing_with_the_body_model_is_rejected() -> None: + upstream = _make_upstream(1) + request = _make_request( + { + "authorization": "Bearer sk-mpkey", + "x-routstr-model-path": encode_model_path( + "http://localhost", 1, "other-model" + ), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + + response = await _run_proxy(request, [(MagicMock(), upstream)]) + + assert response.status_code == 400 + upstream.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_malformed_model_path_header_is_rejected() -> None: + upstream = _make_upstream(1) + request = _make_request( + { + "authorization": "Bearer sk-mpkey", + "x-routstr-model-path": "not-a-model-path", + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + + response = await _run_proxy(request, [(MagicMock(), upstream)]) + + assert response.status_code == 400 + upstream.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_endpoint_tag_pins_the_upstream_subprovider() -> None: + upstream = _make_upstream(1) + upstream.base_url = "https://openrouter.ai/api/v1" + request = _make_request( + { + "authorization": "Bearer sk-mpkey", + "x-routstr-model-path": encode_model_path( + "https://openrouter.ai/api/v1", 1, MODEL_ID, "deepinfra/fp8" + ), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + + await _run_proxy(request, [(MagicMock(), upstream)]) + + forwarded_body = upstream.forward_request.await_args.args[3] + assert json.loads(forwarded_body)["provider"] == { + "order": ["deepinfra/fp8"], + "allow_fallbacks": False, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "raw", + ["", " ", "url=http://localhost&provider-id=1&model-id=test-model&provider-id=2"], +) +async def test_ambiguous_headers_do_not_route(raw: str) -> None: + upstream = _make_upstream(1) + request = _make_request( + {"authorization": "Bearer key", "x-routstr-model-path": raw}, + json.dumps({"model": MODEL_ID}).encode(), + ) + response = await _run_proxy(request, [(MagicMock(), upstream)]) + assert response.status_code == 400 + upstream.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_selector_url_must_match_configured_provider() -> None: + upstream = _make_upstream(1) + request = _make_request( + { + "authorization": "Bearer key", + "x-routstr-model-path": encode_model_path( + "http://169.254.169.254", 1, MODEL_ID + ), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + response = await _run_proxy(request, [(MagicMock(), upstream)]) + assert response.status_code == 404 + upstream.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_endpoint_pin_cannot_be_stripped_on_retry() -> None: + upstream = _make_upstream(1, 400) + upstream.base_url = "https://openrouter.ai/api/v1" + upstream.forward_request.return_value.body = ( + b'{"error":{"message":"provider is not supported"}}' + ) + request = _make_request( + { + "authorization": "Bearer key", + "x-routstr-model-path": encode_model_path( + upstream.base_url, 1, MODEL_ID, "deepinfra/fp8" + ), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + response = await _run_proxy(request, [(MagicMock(), upstream)]) + assert response.status_code == 400 + upstream.forward_request.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "path,handler", + [ + ("v1/chat/completions", "handle_x_cashu"), + ("v1/responses", "handle_x_cashu_responses"), + ], +) +async def test_cashu_receives_endpoint_pinned_body(path: str, handler: str) -> None: + upstream = _make_upstream(1) + upstream.base_url = "https://openrouter.ai/api/v1" + captured: dict[str, Any] = {} + + async def handle(request: Any, *args: Any, **kwargs: Any) -> Any: + captured.update(json.loads(kwargs.get("request_body") or await request.body())) + return MagicMock(status_code=200) + + setattr(upstream, handler, AsyncMock(side_effect=handle)) + request = _make_request( + { + "x-cashu": "test-token", + "x-routstr-model-path": encode_model_path( + upstream.base_url, 1, MODEL_ID, "deepinfra/fp8" + ), + }, + json.dumps( + { + "model": MODEL_ID, + "provider": { + "order": ["other"], + "allow_fallbacks": True, + "data_collection": "deny", + }, + } + ).encode(), + ) + await _run_proxy(request, [(MagicMock(), upstream)], path) + assert captured["provider"] == { + "order": ["deepinfra/fp8"], + "allow_fallbacks": False, + "data_collection": "deny", + } + + +def test_model_path_header_is_not_forwarded() -> None: + from routstr.upstream.base import BaseUpstreamProvider + + upstream = BaseUpstreamProvider("http://localhost", "upstream-key") + headers = upstream.prepare_headers({"x-routstr-model-path": "private-routing-data"}) + assert "x-routstr-model-path" not in headers + + +@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: + import httpx + from fastapi.responses import Response + + from routstr.upstream.openrouter import OpenRouterUpstreamProvider + + upstream = OpenRouterUpstreamProvider(api_key="upstream-key") + upstream.db_id = 1 + fallback = _make_upstream(2) + fallback.handle_x_cashu = AsyncMock() + fallback.handle_x_cashu_responses = AsyncMock() + sent: list[httpx.Request] = [] + + def respond(request: httpx.Request) -> httpx.Response: + sent.append(request) + return httpx.Response(status_code, json={"error": "test"}) + + client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + model = MagicMock(id=MODEL_ID, forwarded_model_id=None, canonical_slug=None) + request = _make_request( + { + "x-cashu": "test-token", + "x-routstr-model-path": encode_model_path( + upstream.base_url, 1, MODEL_ID, "deepinfra/fp8" + ), + }, + json.dumps({"model": MODEL_ID, "provider": {"allow_fallbacks": True}}).encode(), + ) + request.query_params = {} + with ( + patch("routstr.upstream.base.httpx.AsyncClient", return_value=client), + patch( + "routstr.upstream.base.recieve_token", + AsyncMock(return_value=(1000, "msat", "https://mint.test")), + ) as redeem, + patch("routstr.upstream.base.store_cashu_transaction", AsyncMock()), + patch.object(upstream, "send_refund", AsyncMock(return_value="refund")), + patch.object( + upstream, + "handle_x_cashu_chat_completion", + AsyncMock(return_value=Response(status_code=200)), + ), + patch.object( + upstream, + "handle_x_cashu_responses_completion", + AsyncMock(return_value=Response(status_code=200)), + ), + ): + response = await _run_proxy( + request, [(model, upstream), (model, fallback)], path + ) + assert response.status_code == status_code + redeem.assert_awaited_once() + assert len(sent) == 1 + assert sent[0].url.host == "openrouter.ai" + assert json.loads(sent[0].content)["provider"] == { + "order": ["deepinfra/fp8"], + "allow_fallbacks": False, + } + assert "x-routstr-model-path" not in sent[0].headers + fallback.handle_x_cashu.assert_not_awaited() + fallback.handle_x_cashu_responses.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_unpinned_requests_still_fall_back() -> None: + first, fallback = _make_upstream(1, 502), _make_upstream(2) + request = _make_request( + {"authorization": "Bearer key"}, json.dumps({"model": MODEL_ID}).encode() + ) + response = await _run_proxy( + request, [(MagicMock(), first), (MagicMock(), fallback)] + ) + assert response.status_code == 200 + first.forward_request.assert_awaited_once() + fallback.forward_request.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_pinned_exception_does_not_fall_back() -> None: + from routstr.core.exceptions import UpstreamError + + first, fallback = _make_upstream(1), _make_upstream(2) + first.forward_request.side_effect = UpstreamError("unavailable", status_code=503) + request = _make_request( + { + "authorization": "Bearer key", + "x-routstr-model-path": encode_model_path(first.base_url, 1, MODEL_ID), + }, + json.dumps({"model": MODEL_ID}).encode(), + ) + response = await _run_proxy( + request, [(MagicMock(), first), (MagicMock(), fallback)] + ) + assert response.status_code == 503 + first.forward_request.assert_awaited_once() + fallback.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "path,is_ehbp", [("v1/messages", False), ("v1/chat/completions", True)] +) +async def test_unsupported_endpoint_pins_fail_before_payment( + path: str, is_ehbp: bool +) -> None: + upstream = _make_upstream(1) + upstream.base_url = "https://openrouter.ai/api/v1" + headers = { + "x-cashu": "test-token", + "x-routstr-model-path": encode_model_path( + upstream.base_url, 1, MODEL_ID, "deepinfra/fp8" + ), + } + if is_ehbp: + headers.update({"ehbp-encapsulated-key": "sealed", "x-routstr-model": MODEL_ID}) + request = _make_request(headers, json.dumps({"model": MODEL_ID}).encode()) + with ( + patch.object(proxy_module, "check_token_balance") as payment, + patch.object( + proxy_module, "get_candidates", return_value=[(MagicMock(), upstream)] + ), + ): + response = await proxy_module.proxy(request, path, MagicMock()) + assert response.status_code == 400 + assert json.loads(response.body)["error"]["type"] == "unsupported_request" + payment.assert_not_called() + upstream.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cashu", [False, True]) +async def test_ehbp_pin_does_not_fall_back(cashu: bool) -> None: + from routstr.core.exceptions import UpstreamError + + selected, fallback = _make_upstream(1), _make_upstream(2) + selected.supports_ehbp = fallback.supports_ehbp = True + headers = { + "ehbp-encapsulated-key": "sealed", + "x-routstr-model": MODEL_ID, + "x-routstr-model-path": encode_model_path(selected.base_url, 1, MODEL_ID), + } + headers.update({"x-cashu": "token"} if cashu else {"authorization": "Bearer key"}) + request = _make_request(headers, b"encrypted-body") + handler = "forward_ehbp_x_cashu_request" if cashu else "forward_ehbp_request" + with patch.object( + proxy_module, + handler, + AsyncMock(side_effect=UpstreamError("unavailable", status_code=503)), + ) as forward: + response = await _run_proxy( + request, [(MagicMock(), selected), (MagicMock(), fallback)] + ) + assert response.status_code == 503 + forward.assert_awaited_once() + assert forward.await_args is not None + assert forward.await_args.kwargs["upstream"] is selected + + +@pytest.mark.asyncio +async def test_duplicate_header_fields_are_rejected() -> None: + from starlette.datastructures import Headers + + selected = _make_upstream(1) + route = encode_model_path(selected.base_url, 1, MODEL_ID).encode() + request = _make_request({}, json.dumps({"model": MODEL_ID}).encode()) + request.headers = Headers( + raw=[ + (b"authorization", b"Bearer key"), + (b"x-routstr-model-path", route), + (b"x-routstr-model-path", route), + ] + ) + response = await _run_proxy(request, [(MagicMock(), selected)]) + assert response.status_code == 400 + selected.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_attestation_does_not_ignore_model_path() -> None: + request = _make_request( + {"x-routstr-model-path": encode_model_path("http://localhost", 1, MODEL_ID)}, + b"", + ) + request.method = "GET" + with patch.object(proxy_module, "_select_unauthenticated_get_upstreams") as select: + response = await _run_proxy(request, [], "attestation") + assert response.status_code == 400 + select.assert_not_called() + + +@pytest.mark.asyncio +async def test_model_fallback_list_is_rejected_when_pinned() -> None: + selected = _make_upstream(1) + request = _make_request( + { + "authorization": "Bearer key", + "x-routstr-model-path": encode_model_path(selected.base_url, 1, MODEL_ID), + }, + json.dumps({"model": MODEL_ID, "models": ["other-model"]}).encode(), + ) + response = await _run_proxy(request, [(MagicMock(), selected)]) + assert response.status_code == 400 + selected.forward_request.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint", [None, "deepinfra/fp8"]) +@pytest.mark.parametrize("path", ["v1/chat/completions", "v1/responses"]) +@pytest.mark.parametrize("final_status", [200, 429]) +async def test_pinned_recovery_stays_on_selected_provider( + endpoint: str | None, path: str, final_status: int +) -> None: + selected, fallback = _make_upstream(1), _make_upstream(2) + selected.base_url = "https://openrouter.ai/api/v1" + handler = ( + "forward_responses_request" if path == "v1/responses" else "forward_request" + ) + forward = AsyncMock( + side_effect=[ + MagicMock( + status_code=400, + body=b'{"error":{"message":"temperature is deprecated"}}', + ), + MagicMock(status_code=final_status, body=b"{}"), + ] + ) + setattr(selected, handler, forward) + setattr(fallback, handler, AsyncMock()) + request = _make_request( + { + "authorization": "Bearer key", + "x-routstr-model-path": encode_model_path( + selected.base_url, 1, MODEL_ID, endpoint + ), + }, + json.dumps( + { + "model": MODEL_ID, + "temperature": 0.7, + "provider": {"data_collection": "deny"}, + } + ).encode(), + ) + + response = await _run_proxy( + request, [(MagicMock(), selected), (MagicMock(), fallback)], path + ) + + assert response.status_code == final_status + assert forward.await_count == 2 + before, after = [json.loads(call.args[3]) for call in forward.await_args_list] + assert "temperature" in before + assert "temperature" not in after + assert after["model"] == before["model"] == MODEL_ID + assert after["provider"] == before["provider"] + if endpoint: + assert after["provider"]["order"] == [endpoint] + assert after["provider"]["allow_fallbacks"] is False + getattr(fallback, handler).assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("field", ["model", "provider"]) +@pytest.mark.parametrize("endpoint", [None, "deepinfra/fp8"]) +async def test_pinned_recovery_preserves_routing_fields( + field: str, endpoint: str | None +) -> None: + selected, fallback = _make_upstream(1, 400), _make_upstream(2) + selected.base_url = "https://openrouter.ai/api/v1" + selected.forward_request.return_value.body = json.dumps( + {"error": {"message": f"{field} is not supported"}} + ).encode() + request = _make_request( + { + "authorization": "Bearer key", + "x-routstr-model-path": encode_model_path( + selected.base_url, 1, MODEL_ID, endpoint + ), + }, + json.dumps( + {"model": MODEL_ID, "provider": {"data_collection": "deny"}} + ).encode(), + ) + + response = await _run_proxy( + request, [(MagicMock(), selected), (MagicMock(), fallback)] + ) + + assert response.status_code == 400 + selected.forward_request.assert_awaited_once() + fallback.forward_request.assert_not_awaited() diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 240a4199..caab67b9 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -28,6 +28,7 @@ os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") os.environ.setdefault("UPSTREAM_API_KEY", "test") from routstr.core.db import ModelRow, UpstreamProviderRow # noqa: E402 +from routstr.payment import price as price_module # noqa: E402 from routstr.payment.models import models_router # noqa: E402 from routstr.upstream import model_paths as mp # noqa: E402 from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 @@ -202,6 +203,75 @@ async def patched_session( await engine.dispose() +# 1 sat = $0.00005, the quote the path pricing converts with in these tests. +_QUOTE = 5.0e-5 +_DEFAULT_FEE = 1.01 + + +@pytest.fixture +def sats_quote(monkeypatch: pytest.MonkeyPatch) -> float: + monkeypatch.setattr(price_module, "sats_usd_price", lambda: _QUOTE) + return _QUOTE + + +def _priced_endpoints_response() -> httpx.Response: + """Two endpoints for one model, priced and sized differently.""" + return httpx.Response( + 200, + json={ + "data": { + "id": "anthropic/claude-opus-4.6", + "name": "Claude Opus 4.6", + "description": "Anthropic's most capable model", + "architecture": { + "input_modalities": ["text", "image"], + "output_modalities": ["text"], + "tokenizer": "Claude", + "instruct_type": None, + }, + "endpoints": [ + { + "provider_name": "Anthropic", + "tag": "anthropic", + "context_length": 200_000, + "pricing": { + "prompt": "0.000005", + "completion": "0.000025", + }, + }, + { + "provider_name": "Google", + "tag": "google-vertex/us", + "context_length": 128_000, + "pricing": { + "prompt": "0.000003", + "completion": "0.000015", + }, + }, + ], + } + }, + ) + + +def _models_by_endpoint(payload: dict, model_id: str) -> dict[str, dict]: + entry = next(item for item in payload["data"] if item["id"] == model_id) + return { + path["endpoint"]["tag"]: path["model"] + for path in entry["paths"] + if path["endpoint"] is not None + } + + +async def _set_provider_fee(engine: AsyncEngine, provider_id: int, fee: float) -> None: + async with AsyncSession(engine) as session: + provider = await session.get(UpstreamProviderRow, provider_id) + assert provider is not None + provider.provider_fee = fee + session.add(provider) + await session.commit() + + def _paths_of(payload: dict, model_id: str) -> set[str]: for entry in payload["data"]: if entry["id"] == model_id: @@ -244,6 +314,12 @@ def _path_entry( or ("anthropic" if provider_id == 1 else "openrouter"), }, "endpoint": endpoint, + "model": { + "id": model_id, + "forwarded_model_id": None, + "canonical_slug": None, + "enabled": True, + }, } @@ -252,10 +328,22 @@ def _path_entry( # --------------------------------------------------------------------------- # -def test_is_openrouter_base_url() -> None: - assert mp.is_openrouter_base_url("https://openrouter.ai/api/v1") is True - assert mp.is_openrouter_base_url("https://api.anthropic.com") is False - assert mp.is_openrouter_base_url(None) is False +@pytest.mark.parametrize( + "url, expected", + [ + ("https://openrouter.ai/api/v1", True), + ("https://OPENROUTER.AI/api/v1", True), + ("https://api.anthropic.com", False), + ("https://openrouter.ai.evil.test/api/v1", False), + ("https://evil.test/openrouter.ai", False), + ("https://openrouter.ai@evil.test/api/v1", False), + ("https://[invalid", False), + ("", False), + (None, False), + ], +) +def test_is_openrouter_base_url(url: str | None, expected: bool) -> None: + assert mp.is_openrouter_base_url(url) is expected def test_native_anthropic_not_openrouter() -> None: @@ -367,6 +455,37 @@ async def test_direct_provider_single_path_uses_provider_type( assert payload["updated_at"] is not None +@pytest.mark.asyncio +async def test_get_all_model_paths_includes_details_for_each_path( + patched_session: AsyncEngine, +) -> None: + model = _model("claude-opus-4.6") + model.name = "Claude Opus 4.6" + model.description = "Anthropic's most capable model" + model.pricing = {"prompt": 0.000001, "completion": 0.000002} + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[model], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + payload = await mp.get_all_model_paths() + + expected_path = _path_entry(1, "claude-opus-4.6") + expected_path["model"] = { + "id": "claude-opus-4.6", + "forwarded_model_id": None, + "canonical_slug": None, + "enabled": True, + "name": "Claude Opus 4.6", + "description": "Anthropic's most capable model", + "pricing": {"prompt": 0.000001, "completion": 0.000002}, + } + assert payload["data"] == [{"id": "claude-opus-4.6", "paths": [expected_path]}] + + @pytest.mark.asyncio async def test_direct_path_masks_private_configured_provider_url( patched_session: AsyncEngine, @@ -564,9 +683,12 @@ async def test_refresh_model_paths_uses_db_forwarded_alias( await mp.refresh_model_paths([provider]) - assert (await mp.get_all_model_paths())["data"] == [ - {"id": "public-alias", "paths": [_path_entry(1, "public-alias")]} - ] + payload = await mp.get_all_model_paths() + assert _paths_of(payload, "public-alias") == {_expected_path(1, "public-alias")} + model = payload["data"][0]["paths"][0]["model"] + assert model["id"] == "public-alias" + assert model["description"] == "test model" + assert model["pricing"] == {"prompt": 0.000001, "completion": 0.000002} @pytest.mark.asyncio @@ -586,12 +708,14 @@ async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cac await mp.refresh_model_paths([provider]) - assert (await mp.get_all_model_paths())["data"] == [ - { - "id": "public-deployment", - "paths": [_path_entry(1, "public-deployment")], - } - ] + payload = await mp.get_all_model_paths() + assert _paths_of(payload, "public-deployment") == { + _expected_path(1, "public-deployment") + } + model = payload["data"][0]["paths"][0]["model"] + assert model["id"] == "public-deployment" + assert model["description"] == "test model" + assert model["context_length"] == 8192 @pytest.mark.asyncio @@ -765,6 +889,148 @@ async def test_openrouter_provider_adds_endpoint_paths( assert {item["provider"]["id"] for item in payload["data"]} == {2} +@pytest.mark.asyncio +async def test_openrouter_paths_include_endpoint_specific_model_prices( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch, sats_quote: float +) -> None: + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _priced_endpoints_response()) + + await mp.refresh_model_paths([provider]) + + payload = await mp.get_all_model_paths() + assert payload["data"][0]["id"] == "claude-opus-4.6" + models = _models_by_endpoint(payload, "claude-opus-4.6") + + anthropic = models["anthropic"] + google = models["google-vertex/us"] + assert anthropic["description"] == "Anthropic's most capable model" + assert google["description"] == "Anthropic's most capable model" + assert anthropic["pricing"]["prompt"] == pytest.approx(0.000005 * _DEFAULT_FEE) + assert anthropic["pricing"]["completion"] == pytest.approx(0.000025 * _DEFAULT_FEE) + assert google["pricing"]["prompt"] == pytest.approx(0.000003 * _DEFAULT_FEE) + assert google["pricing"]["completion"] == pytest.approx(0.000015 * _DEFAULT_FEE) + assert anthropic["context_length"] == 200_000 + assert google["context_length"] == 128_000 + + +@pytest.mark.asyncio +async def test_endpoint_paths_are_priced_in_sats_from_their_own_rates( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch, sats_quote: float +) -> None: + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _priced_endpoints_response()) + + await mp.refresh_model_paths([provider]) + + models = _models_by_endpoint(await mp.get_all_model_paths(), "claude-opus-4.6") + anthropic = models["anthropic"]["sats_pricing"] + google = models["google-vertex/us"]["sats_pricing"] + + assert anthropic["prompt"] == pytest.approx(0.000005 * _DEFAULT_FEE / sats_quote) + assert anthropic["completion"] == pytest.approx( + 0.000025 * _DEFAULT_FEE / sats_quote + ) + assert google["prompt"] == pytest.approx(0.000003 * _DEFAULT_FEE / sats_quote) + assert google["completion"] == pytest.approx(0.000015 * _DEFAULT_FEE / sats_quote) + + +@pytest.mark.asyncio +async def test_endpoint_max_costs_use_that_endpoint_context_length( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch, sats_quote: float +) -> None: + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _priced_endpoints_response()) + + await mp.refresh_model_paths([provider]) + + models = _models_by_endpoint(await mp.get_all_model_paths(), "claude-opus-4.6") + # Max cost is the context window billed at the dearer of the two rates. + assert models["anthropic"]["sats_pricing"]["max_cost"] == pytest.approx( + 200_000 * 0.000025 * _DEFAULT_FEE / sats_quote + ) + assert models["google-vertex/us"]["sats_pricing"]["max_cost"] == pytest.approx( + 128_000 * 0.000015 * _DEFAULT_FEE / sats_quote + ) + + +@pytest.mark.asyncio +async def test_path_pricing_uses_the_provider_fee_of_its_own_provider( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch, sats_quote: float +) -> None: + await _set_provider_fee(patched_session, 2, 1.5) + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _priced_endpoints_response()) + + await mp.refresh_model_paths([provider]) + + models = _models_by_endpoint(await mp.get_all_model_paths(), "claude-opus-4.6") + anthropic = models["anthropic"] + assert anthropic["pricing"]["prompt"] == pytest.approx(0.000005 * 1.5) + assert anthropic["sats_pricing"]["prompt"] == pytest.approx( + 0.000005 * 1.5 / sats_quote + ) + + +@pytest.mark.asyncio +async def test_paths_keep_upstream_pricing_when_the_quote_is_unavailable( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + def _no_quote() -> float: + raise ValueError("SATS price not initialized") + + monkeypatch.setattr(price_module, "sats_usd_price", _no_quote) + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _priced_endpoints_response()) + + await mp.refresh_model_paths([provider]) + + models = _models_by_endpoint(await mp.get_all_model_paths(), "claude-opus-4.6") + anthropic = models["anthropic"] + assert "sats_pricing" not in anthropic + assert anthropic["pricing"] == { + "prompt": "0.000005", + "completion": "0.000025", + } + + +@pytest.mark.asyncio +async def test_already_priced_metadata_is_not_priced_again( + patched_session: AsyncEngine, sats_quote: float +) -> None: + model = _model("claude-opus-4.6") + model.pricing = {"prompt": 0.000001, "completion": 0.000002} + model.sats_pricing = {"prompt": 0.02, "completion": 0.04} + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[model], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + + payload = await mp.get_all_model_paths() + priced = payload["data"][0]["paths"][0]["model"] + assert priced["sats_pricing"] == {"prompt": 0.02, "completion": 0.04} + assert priced["pricing"] == {"prompt": 0.000001, "completion": 0.000002} + + @pytest.mark.asyncio async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch diff --git a/tests/unit/test_nostr_listing_announce.py b/tests/unit/test_nostr_listing_announce.py new file mode 100644 index 00000000..69bc578e --- /dev/null +++ b/tests/unit/test_nostr_listing_announce.py @@ -0,0 +1,176 @@ +"""Tests for the kind 38421 provider announcement loop. + +The regression these cover: the node runner configures the NSEC through the +admin UI (not the ``.env`` file), which only mutates the live ``settings`` +singleton. The announcement task must therefore (a) already be running, and +(b) pick the new identity up on its own — without a process restart. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +from routstr.nostr import listing + +NSEC_A = "11" * 32 +NSEC_B = "22" * 32 + + +def _quiet_settings(monkeypatch: Any, nsec: str = "") -> None: + """Pin the settings the loop reads, and keep it off the network/Tor.""" + monkeypatch.setattr(listing.settings, "nsec", nsec) + monkeypatch.setattr(listing.settings, "http_url", "https://node.example.com") + monkeypatch.setattr(listing.settings, "onion_url", "") + monkeypatch.setattr(listing.settings, "relays", []) + monkeypatch.setattr(listing.settings, "provider_id", "testprovider") + monkeypatch.setattr(listing.settings, "cashu_mints", []) + monkeypatch.setattr(listing, "discover_onion_url_from_tor", lambda: None) + + +def _capture_publishes(monkeypatch: Any) -> list[dict[str, Any]]: + published: list[dict[str, Any]] = [] + + async def fake_query( + *args: Any, **kwargs: Any + ) -> tuple[list[dict[str, Any]], bool]: + return [], True + + async def fake_publish( + relay_url: str, event: dict[str, Any], timeout: int = 30 + ) -> bool: + published.append(event) + return True + + monkeypatch.setattr(listing, "query_listing_events", fake_query) + monkeypatch.setattr(listing, "publish_to_relay", fake_publish) + return published + + +def _distinct_events(published: list[dict[str, Any]]) -> list[dict[str, Any]]: + """One entry per announcement; the publisher is invoked once per relay.""" + by_id: dict[str, dict[str, Any]] = {} + for event in published: + by_id[event["id"]] = event + return list(by_id.values()) + + +@pytest.mark.asyncio +async def test_announce_provider_idles_without_nsec_then_publishes_when_saved( + monkeypatch: Any, +) -> None: + """The reported bug: NSEC saved through the admin UI must get announced.""" + sleeps: list[float] = [] + published = _capture_publishes(monkeypatch) + _quiet_settings(monkeypatch, nsec="") + + async def fake_sleep(seconds: float) -> None: + sleeps.append(seconds) + if len(sleeps) == 1: + # The node runner hits "save" in the admin UI: the nsec appears on + # the live settings singleton while the task is already idling. + monkeypatch.setattr(listing.settings, "nsec", NSEC_A) + return + raise asyncio.CancelledError() + + monkeypatch.setattr(listing.asyncio, "sleep", fake_sleep) + + await listing.announce_provider() + + # First pass idles (no nsec), then the announcing pass sleeps one poll tick + # into the re-announce interval. + assert sleeps == [ + listing.DISABLED_POLL_SECONDS, + listing.IDENTITY_POLL_SECONDS, + ] + announcements = _distinct_events(published) + assert len(announcements) == 1 + assert announcements[0]["kind"] == 38421 + + +@pytest.mark.asyncio +async def test_announce_provider_reannounces_when_nsec_is_replaced( + monkeypatch: Any, +) -> None: + """Replacing the key must re-resolve the identity and announce again.""" + sleeps: list[float] = [] + published = _capture_publishes(monkeypatch) + _quiet_settings(monkeypatch, nsec=NSEC_A) + + async def fake_sleep(seconds: float) -> None: + sleeps.append(seconds) + if len(sleeps) == 1: + monkeypatch.setattr(listing.settings, "nsec", NSEC_B) + return + raise asyncio.CancelledError() + + monkeypatch.setattr(listing.asyncio, "sleep", fake_sleep) + + await listing.announce_provider() + + assert len(published) == 2 * len(listing.DEFAULT_RELAY_URLS) + announcements = _distinct_events(published) + assert len(announcements) == 2 + pubkeys = {event["pubkey"] for event in announcements} + assert len(pubkeys) == 2, "each identity must be announced under its own pubkey" + + +@pytest.mark.asyncio +async def test_announce_provider_idles_without_endpoints(monkeypatch: Any) -> None: + """No publishable endpoint: idle instead of exiting, and never publish.""" + sleeps: list[float] = [] + published = _capture_publishes(monkeypatch) + _quiet_settings(monkeypatch, nsec=NSEC_A) + monkeypatch.setattr(listing.settings, "http_url", "") + + async def fake_sleep(seconds: float) -> None: + sleeps.append(seconds) + raise asyncio.CancelledError() + + monkeypatch.setattr(listing.asyncio, "sleep", fake_sleep) + + await listing.announce_provider() + + assert published == [] + assert sleeps == [listing.DISABLED_POLL_SECONDS] + + +@pytest.mark.asyncio +async def test_sleep_until_next_announcement_wakes_on_nsec_change( + monkeypatch: Any, +) -> None: + sleeps: list[float] = [] + monkeypatch.setattr(listing.settings, "nsec", NSEC_A) + + async def fake_sleep(seconds: float) -> None: + sleeps.append(seconds) + monkeypatch.setattr(listing.settings, "nsec", NSEC_B) + + monkeypatch.setattr(listing.asyncio, "sleep", fake_sleep) + + await listing._sleep_until_next_announcement( + listing.ANNOUNCEMENT_INTERVAL_SECONDS, NSEC_A + ) + + # Woke on the first poll tick rather than sleeping out the whole interval. + assert sleeps == [listing.IDENTITY_POLL_SECONDS] + + +@pytest.mark.asyncio +async def test_sleep_until_next_announcement_buckets_a_short_interval( + monkeypatch: Any, +) -> None: + sleeps: list[float] = [] + monkeypatch.setattr(listing.settings, "nsec", NSEC_A) + + async def fake_sleep(seconds: float) -> None: + sleeps.append(seconds) + + monkeypatch.setattr(listing.asyncio, "sleep", fake_sleep) + + await listing._sleep_until_next_announcement(45, NSEC_A) + + # 45s interval: a full 30s tick, then the 15s remainder (never overshoots). + assert sleeps == [listing.IDENTITY_POLL_SECONDS, 15] diff --git a/tests/unit/test_nostr_sdk.py b/tests/unit/test_nostr_sdk.py new file mode 100644 index 00000000..ae4a42c0 --- /dev/null +++ b/tests/unit/test_nostr_sdk.py @@ -0,0 +1,129 @@ +from __future__ import annotations + +import json +from typing import Any + +import pytest +from nostr_sdk import Event, SendEventOutput + +from routstr.nostr import sdk +from routstr.nostr.listing import create_listing_event, nsec_to_keypair + +PRIVATE_KEY_HEX = "11" * 32 +PRIVATE_KEY_NSEC = "nsec1zyg3zyg3zyg3zyg3zyg3zyg3zyg3zyg3zyg3zyg3zyg3zyg3zygs4rm7hz" +PUBLIC_KEY_HEX = "4f355bdcb7cc0af728ef3cceb9615d90684bb5b2ca5f859ab0f0b704075871aa" + + +def test_nsec_and_hex_parse_to_same_keypair() -> None: + expected = (PRIVATE_KEY_HEX, PUBLIC_KEY_HEX) + + assert nsec_to_keypair(PRIVATE_KEY_HEX) == expected + assert nsec_to_keypair(PRIVATE_KEY_NSEC) == expected + + +def test_listing_event_is_valid_nip01_event() -> None: + event = create_listing_event( + PRIVATE_KEY_HEX, + "provider123", + ["https://provider.example.com"], + mint_urls=["https://mint.example.com"], + version="1.2.3", + metadata={"name": "Provider"}, + ) + + assert event["pubkey"] == PUBLIC_KEY_HEX + assert event["kind"] == 38421 + assert Event.from_json(json.dumps(event)).verify() + + +class FakeClient: + def __init__( + self, + event: dict[str, Any] | None = None, + send_failure: str | None = None, + ) -> None: + self.event = event + self.send_failure = send_failure + self.relay: Any = None + self.connected = False + self.shutdown_called = False + self.sent_event: Event | None = None + + async def add_relay(self, relay: Any) -> bool: + self.relay = relay + return True + + async def connect(self) -> None: + self.connected = True + + async def fetch_events(self, *args: Any, **kwargs: Any) -> list[Event]: + assert self.event is not None + return [Event.from_json(json.dumps(self.event))] + + async def send_event(self, event: Event, **kwargs: Any) -> SendEventOutput: + self.sent_event = event + if self.send_failure is not None: + return SendEventOutput( + id=event.id(), success=[], failed={self.relay: self.send_failure} + ) + return SendEventOutput(id=event.id(), success=[self.relay], failed={}) + + async def shutdown(self) -> None: + self.shutdown_called = True + + +@pytest.mark.asyncio +async def test_fetch_events_uses_sdk_client_and_closes_it(monkeypatch: Any) -> None: + event = create_listing_event( + PRIVATE_KEY_HEX, + "provider123", + ["https://provider.example.com"], + ) + client = FakeClient(event) + monkeypatch.setattr(sdk, "Client", lambda: client) + + fetched = await sdk.fetch_events( + "wss://relay.example.com", + kind=38421, + author=PUBLIC_KEY_HEX, + limit=10, + timeout=30, + ) + + assert fetched == [event] + assert client.connected + assert client.shutdown_called + + +@pytest.mark.asyncio +async def test_send_event_uses_sdk_client_and_closes_it(monkeypatch: Any) -> None: + event = create_listing_event( + PRIVATE_KEY_HEX, + "provider123", + ["https://provider.example.com"], + ) + client = FakeClient() + monkeypatch.setattr(sdk, "Client", lambda: client) + + await sdk.send_event("wss://relay.example.com", event, timeout=30) + + assert client.connected + assert client.sent_event is not None + assert client.sent_event.verify() + assert client.shutdown_called + + +@pytest.mark.asyncio +async def test_send_event_raises_when_relay_rejects(monkeypatch: Any) -> None: + event = create_listing_event( + PRIVATE_KEY_HEX, + "provider123", + ["https://provider.example.com"], + ) + client = FakeClient(send_failure="blocked: rate limited") + monkeypatch.setattr(sdk, "Client", lambda: client) + + with pytest.raises(RuntimeError, match="rate limited"): + await sdk.send_event("wss://relay.example.com", event, timeout=30) + + assert client.shutdown_called diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index 6809d94c..343d9cdc 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -155,3 +155,652 @@ async def test_discounted_max_cost_floors_at_min_request_msat() -> None: cost = await calculate_discounted_max_cost(150_000, body, model_obj) assert cost == 1000 + + +def test_estimate_prompt_tokens_counts_every_string_in_the_body() -> None: + from routstr.payment.helpers import estimate_prompt_tokens, estimate_tokens + + hidden = "x" * 3_000 # ~1000 tokens of prompt hidden from the text estimator + body: dict[str, Any] = { + "messages": [{"role": "user", "content": "hi"}], + "tools": [ + { + "type": "function", + "function": { + "name": "f", + "description": hidden, + "parameters": {"type": "object", "properties": {hidden: {}}}, + }, + } + ], + } + + # The text-only estimator sees almost nothing; the conservative one sees it. + assert estimate_tokens(body["messages"]) < 10 + assert estimate_prompt_tokens(body) >= 1_000 + + # No carve-out is exempt: neither a caller-chosen key name nor a caller-chosen + # value prefix can buy a discount, so both still count in full. + assert estimate_prompt_tokens({"tools": [{"data": hidden}]}) >= 1_000 + assert estimate_prompt_tokens({"system": "data:" + hidden}) >= 1_000 + assert estimate_prompt_tokens({"prompt": [[1, 2], [3, 4, 5]]}) >= 5 + + +async def test_discount_counts_legacy_token_id_prompt() -> None: + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.0 + pricing.max_prompt_cost = 50.0 + pricing.max_completion_cost = 0.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + body = {"model": "test-model", "prompt": list(range(50_000)), "max_tokens": 0} + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost = await calculate_discounted_max_cost(50_000, body, model_obj) + + assert cost == 50_000 + + +async def test_discount_cannot_be_dodged_by_hiding_prompt_in_tools() -> None: + """A large prompt moved from messages into tool schemas must reserve the + same cost — otherwise a caller undercharges by hiding weight from the + estimator.""" + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.5 + pricing.completion = 0.01 + pricing.max_prompt_cost = 100.0 + pricing.max_completion_cost = 100.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + big_text = "word " * 2_000 + base = {"model": "test-model", "max_tokens": 10} + in_messages = { + **base, + "messages": [{"role": "user", "content": big_text}], + } + hiding_places = { + "tools": { + **base, + "messages": [{"role": "user", "content": "hi"}], + "tools": [ + {"type": "function", "function": {"name": "f", "description": big_text}} + ], + }, + # Anthropic forwards a top-level system prompt; it is billed like any other. + "system": { + **base, + "messages": [{"role": "user", "content": "hi"}], + "system": big_text, + }, + # A key named like an image field must not win an image exclusion. + "image-named key": { + **base, + "messages": [{"role": "user", "content": "hi"}], + "tools": [{"function": {"parameters": {"data": big_text}}}], + }, + # Nor may a caller-chosen "data:" prefix, in any field the body allows. + "data-prefixed content": { + **base, + "messages": [{"role": "user", "content": "data:" + big_text}], + }, + "data-prefixed text block": { + **base, + "messages": [ + { + "role": "user", + "content": [{"type": "text", "text": "data:" + big_text}], + } + ], + }, + "data-prefixed system": { + **base, + "messages": [{"role": "user", "content": "hi"}], + "system": "data:" + big_text, + }, + } + + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost_messages = await calculate_discounted_max_cost( + 150_000, in_messages, model_obj + ) + for where, body in hiding_places.items(): + cost = await calculate_discounted_max_cost(150_000, body, model_obj) + # Same prompt weight → at least the same reservation, never the floor. + assert cost >= cost_messages, where + assert cost > 1000, where + + +async def test_discounted_max_cost_counts_responses_input_images() -> None: + import base64 + from io import BytesIO + + from PIL import Image + + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.001 + pricing.max_prompt_cost = 100.0 + pricing.max_completion_cost = 0.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + image = Image.new("RGB", (512, 512), "red") + buffer = BytesIO() + image.save(buffer, format="JPEG") + data_url = "data:image/jpeg;base64," + base64.b64encode(buffer.getvalue()).decode() + + no_image = { + "model": "test-model", + "input": [{"role": "user", "content": "hi"}], + } + with_image = { + "model": "test-model", + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "hi"}, + {"type": "input_image", "image_url": data_url, "detail": "high"}, + ], + } + ], + } + + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost_no_image = await calculate_discounted_max_cost( + 100_000, no_image, model_obj + ) + cost_with_image = await calculate_discounted_max_cost( + 100_000, with_image, model_obj + ) + + # The 512x512 high-detail image (85 + 170 = 255 tokens) is billed as prompt + # weight, so it reserves strictly more than the identical text-only body. + assert cost_with_image > cost_no_image + + +def _responses_image(url: str, detail: str | None = "original") -> list[dict[str, Any]]: + return [ + { + "role": "user", + "content": [{"type": "input_image", "image_url": url, "detail": detail}], + } + ] + + +async def _responses_image_tokens(input_data: list[dict[str, Any]]) -> int: + from routstr.payment.helpers import estimate_image_tokens_in_messages + from routstr.payment.responses_input import responses_input_to_messages + + messages = responses_input_to_messages(input_data) + assert messages is not None + return await estimate_image_tokens_in_messages(messages) + + +async def test_remote_original_image_is_fetched_not_worst_cased() -> None: + from io import BytesIO + + from PIL import Image + + image = Image.new("RGB", (512, 512), "red") + buffer = BytesIO() + image.save(buffer, format="JPEG") + image_bytes = buffer.getvalue() + + with patch( + "routstr.payment.helpers._fetch_image_from_url", + new=AsyncMock(return_value=image_bytes), + ): + # 256 patches * 1.2 + assert ( + await _responses_image_tokens(_responses_image("https://x.test/i.jpg")) + == 308 + ) + + with patch( + "routstr.payment.helpers._fetch_image_from_url", + new=AsyncMock(return_value=None), + ): + assert ( + await _responses_image_tokens(_responses_image("https://x.test/i.jpg")) + == 36_000 + ) + + +async def test_broken_data_url_reserves_declared_detail_worst_case() -> None: + from routstr.payment.helpers import estimate_image_tokens_in_messages + + broken = "data:image/jpeg;base64,!!!" + assert await _responses_image_tokens(_responses_image(broken)) == 36_000 + assert await _responses_image_tokens(_responses_image(broken, "high")) == 85 + ( + 170 * 4 + ) + + chat = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": broken, "detail": "original"}, + } + ], + } + ] + assert await estimate_image_tokens_in_messages(chat) == 36_000 + + +async def test_chat_original_image_fetch_failure_reserves_worst_case() -> None: + from routstr.payment.helpers import estimate_image_tokens_in_messages + + chat = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": "https://x.test/i.jpg", "detail": "original"}, + } + ], + } + ] + with patch( + "routstr.payment.helpers._fetch_image_from_url", + new=AsyncMock(return_value=None), + ): + assert await estimate_image_tokens_in_messages(chat) == 36_000 + + +async def test_responses_images_share_per_request_fetch_cap() -> None: + from routstr.payment.helpers import IMAGE_FETCH_MAX_PER_REQUEST + + input_data = [ + { + "role": "user", + "content": [ + { + "type": "input_image", + "image_url": f"https://x.test/{i}.jpg", + "detail": "original", + } + for i in range(IMAGE_FETCH_MAX_PER_REQUEST + 1) + ], + } + ] + fetch = AsyncMock(return_value=None) + with patch("routstr.payment.helpers._fetch_image_from_url", new=fetch): + tokens = await _responses_image_tokens(input_data) + + assert fetch.await_count == IMAGE_FETCH_MAX_PER_REQUEST + assert tokens == 36_000 * (IMAGE_FETCH_MAX_PER_REQUEST + 1) + + +async def test_responses_transform_failure_falls_back_to_worst_case() -> None: + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.001 + pricing.max_prompt_cost = 100.0 + pricing.max_completion_cost = 0.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + body = { + "model": "test-model", + "input": [ + { + "role": "user", + "content": [ + {"type": "input_image", "image_url": "https://x.test/a.jpg"}, + {"type": "input_image", "image_url": "https://x.test/b.jpg"}, + ], + } + ], + } + fetch = AsyncMock(return_value=None) + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + patch("routstr.payment.helpers._fetch_image_from_url", new=fetch), + patch( + "routstr.payment.responses_input.LiteLLMCompletionResponsesConfig." + "transform_responses_api_input_to_messages", + side_effect=RuntimeError("boom"), + ), + ): + cost = await calculate_discounted_max_cost(100_000, body, model_obj) + + fetch.assert_not_awaited() + # 2 * 36,000 tokens * 0.001 sats = 72 sats reserved + assert 72_000 <= cost < 100_000 + + +async def test_discounted_max_cost_body_max_output_tokens_fallback() -> None: + """Body ``max_output_tokens`` (Responses API) is honored as a completion cap.""" + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.001 + pricing.max_prompt_cost = 0.0 + pricing.max_completion_cost = 100.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + body = {"max_output_tokens": 80_000} + + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost = await calculate_discounted_max_cost(100_000, body, model_obj) + + assert cost == 80_000 + + +async def test_discounted_max_cost_body_max_completion_tokens_fallback() -> None: + """Body ``max_completion_tokens`` (modern chat) is honored as a completion cap.""" + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.001 + pricing.max_prompt_cost = 0.0 + pricing.max_completion_cost = 100.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + body = {"max_completion_tokens": 80_000} + + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost = await calculate_discounted_max_cost(100_000, body, model_obj) + + assert cost == 80_000 + + +async def test_discounted_max_cost_uses_largest_completion_cap() -> None: + """With several completion caps declared, the largest bounds the reservation. + + Upstream precedence between ``max_tokens`` / ``max_completion_tokens`` / + ``max_output_tokens`` varies by provider, so reserving against anything + but the largest could under-cover what the upstream bills. + """ + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.001 + pricing.max_prompt_cost = 0.0 + pricing.max_completion_cost = 100.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + body = { + "max_tokens": 50_000, + "max_completion_tokens": 10_000, + "max_output_tokens": 80_000, + } + + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost = await calculate_discounted_max_cost(100_000, body, model_obj) + + # 80_000 is the largest declared cap: 100.0 - 80.0 = 20 sats discount. + assert cost == 80_000 + + +async def test_discounted_max_cost_invalid_completion_cap_ignored() -> None: + """Unparseable caps yield no completion discount rather than under-reserving.""" + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.001 + pricing.max_prompt_cost = 0.0 + pricing.max_completion_cost = 100.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + # No valid cap at all -> no completion discount. + cost = await calculate_discounted_max_cost( + 100_000, {"max_completion_tokens": "sixty-four-k"}, model_obj + ) + assert cost == 100_000 + + # An invalid sibling does not poison a valid cap on another field. + cost = await calculate_discounted_max_cost( + 100_000, + {"max_tokens": "bad", "max_completion_tokens": 80_000}, + model_obj, + ) + assert cost == 80_000 + + +def _responses_file_image(detail: str | None) -> list[dict[str, Any]]: + part: dict[str, Any] = {"type": "input_image", "file_id": "file-1"} + if detail is not None: + part["detail"] = detail + return [{"role": "user", "content": [part]}] + + +async def test_responses_input_detail_and_file_id() -> None: + import base64 + from io import BytesIO + + from PIL import Image + + # file_id: dimensions can't be fetched, so use a conservative max-size + # estimate (4 tiles for auto/high) and honor the detail sibling for low. + fetch = AsyncMock(return_value=None) + with patch("routstr.payment.helpers._fetch_image_from_url", new=fetch): + assert await _responses_image_tokens(_responses_file_image(None)) == 85 + ( + 170 * 4 + ) + assert await _responses_image_tokens(_responses_file_image("low")) == 85 + fetch.assert_not_awaited() + + # image_url honors the sibling detail instead of always defaulting to auto. + image = Image.new("RGB", (512, 512), "red") + buffer = BytesIO() + image.save(buffer, format="JPEG") + data_url = "data:image/jpeg;base64," + base64.b64encode(buffer.getvalue()).decode() + + assert await _responses_image_tokens(_responses_image(data_url, "low")) == 85 + assert ( + await _responses_image_tokens(_responses_image(data_url, "high")) == 85 + 170 + ) # 512x512 = 1 tile + + +def test_calculate_image_tokens_original_detail() -> None: + from routstr.payment.helpers import _calculate_image_tokens + + # Patch-based pricing: ceil(patches * 1.2) tokens at 32x32px patches. + assert _calculate_image_tokens(640, 640, "original") == 480 # 400 patches + assert _calculate_image_tokens(2048, 2048, "original") == 4_916 # 4,096 patches + # The same image on the tiled high-detail path caps at 765 tokens. + assert _calculate_image_tokens(2048, 2048, "high") == 765 + # Above the 30,000-patch rejection limit the estimate is capped at + # 36,000 tokens (30,000 patches * 1.2). + assert _calculate_image_tokens(10_000, 10_000, "original") == 36_000 + + +async def test_responses_input_original_detail() -> None: + import base64 + from io import BytesIO + + from PIL import Image + + image = Image.new("RGB", (2048, 2048), "red") + buffer = BytesIO() + image.save(buffer, format="JPEG") + data_url = "data:image/jpeg;base64," + base64.b64encode(buffer.getvalue()).decode() + + # image_url: billed at the decoded original resolution (4,096 patches), + # not the 765-token tile cap. + assert await _responses_image_tokens(_responses_image(data_url)) == 4_916 + + # file_id: dimensions unknown, so use the 30,000-patch worst case. + assert await _responses_image_tokens(_responses_file_image("original")) == 36_000 + + # Explicit null detail behaves like the auto default (tiled math). + assert await _responses_image_tokens(_responses_image(data_url, None)) == 85 + ( + 170 * 4 + ) + + +def test_responses_input_to_messages_shapes() -> None: + from routstr.payment.responses_input import ( + FILE_ID_URL_PREFIX, + count_input_images, + responses_input_to_messages, + ) + + input_data = [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "hi"}, + {"type": "input_image", "file_id": "file-1", "detail": "original"}, + ], + }, + {"type": "function_call_output", "call_id": "c1", "output": "out"}, + ] + messages = responses_input_to_messages(input_data) + assert messages is not None + assert messages[0]["role"] == "user" + parts = messages[0]["content"] + assert parts[0] == {"type": "text", "text": "hi"} + assert parts[1]["type"] == "image_url" + assert parts[1]["image_url"] == { + "url": f"{FILE_ID_URL_PREFIX}file-1", + "detail": "original", + } + assert messages[1]["role"] == "tool" + + # dict-form image_url: litellm nests it verbatim, so it is flattened first. + nested = responses_input_to_messages( + [ + { + "type": "message", + "role": "user", + "content": [ + { + "type": "input_image", + "image_url": { + "url": "https://x.test/a.jpg", + "detail": "original", + }, + } + ], + } + ] + ) + assert nested is not None + assert nested[0]["content"][0]["image_url"] == { + "url": "https://x.test/a.jpg", + "detail": "original", + } + + assert responses_input_to_messages("plain") == [ + {"role": "user", "content": "plain"} + ] + assert responses_input_to_messages(None) == [] + assert count_input_images(input_data) == 1 + + +async def test_estimate_image_tokens_in_messages_original_detail() -> None: + """Chat Completions also accepts original detail via the nested dict.""" + import base64 + from io import BytesIO + + from PIL import Image + + from routstr.payment.helpers import estimate_image_tokens_in_messages + + image = Image.new("RGB", (640, 640), "blue") + buffer = BytesIO() + image.save(buffer, format="JPEG") + data_url = "data:image/jpeg;base64," + base64.b64encode(buffer.getvalue()).decode() + + messages = [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": data_url, "detail": "original"}, + } + ], + } + ] + # 640x640 -> 20x20 = 400 patches -> ceil(400 * 1.2) = 480 tokens. + assert await estimate_image_tokens_in_messages(messages) == 480 + + messages = [ + { + "role": "user", + "content": [ + {"type": "input_image", "file_id": "file-1", "detail": "original"} + ], + } + ] + assert await estimate_image_tokens_in_messages(messages) == 36_000 diff --git a/tests/unit/test_ppq_resilience.py b/tests/unit/test_ppq_resilience.py new file mode 100644 index 00000000..b1dd9f9c --- /dev/null +++ b/tests/unit/test_ppq_resilience.py @@ -0,0 +1,103 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from routstr.upstream.ppqai import ( + PPQAIUpstreamProvider, + PPQCircuitOpenError, + _ppq_circuits, + _safe_read_request, +) + + +@pytest.fixture(autouse=True) +def _clear_ppq_circuits() -> None: + _ppq_circuits.clear() + + +@pytest.mark.asyncio +async def test_safe_ppq_read_retries_timeout_with_bounded_backoff() -> None: + request = httpx.Request("GET", "https://api.ppq.ai/models") + success = httpx.Response(200, request=request, json={"data": []}) + client = MagicMock() + client.request = AsyncMock( + side_effect=[httpx.ReadTimeout("", request=request), success] + ) + with ( + patch("routstr.upstream.ppqai.random.uniform", return_value=0.0), + patch("routstr.upstream.ppqai.asyncio.sleep", AsyncMock()) as sleep, + ): + response = await _safe_read_request( + client, "GET", "https://api.ppq.ai/models", headers={} + ) + + assert response is success + assert client.request.await_count == 2 + sleep.assert_awaited_once_with(0.25) + + +@pytest.mark.asyncio +async def test_safe_ppq_read_opens_cross_cycle_circuit_and_probe_clears_it() -> None: + request = httpx.Request("GET", "https://api.ppq.ai/models") + client = MagicMock() + client.request = AsyncMock(side_effect=httpx.ReadTimeout("down", request=request)) + + with ( + patch("routstr.upstream.ppqai.random.uniform", return_value=0.0), + patch("routstr.upstream.ppqai.asyncio.sleep", AsyncMock()), + patch("routstr.upstream.ppqai.time.monotonic", return_value=100.0), + pytest.raises(httpx.ReadTimeout), + ): + await _safe_read_request(client, "GET", str(request.url), headers={}) + assert client.request.await_count == 3 + + with ( + patch("routstr.upstream.ppqai.time.monotonic", return_value=110.0), + pytest.raises(PPQCircuitOpenError), + ): + await _safe_read_request(client, "GET", str(request.url), headers={}) + assert client.request.await_count == 3 + + success = httpx.Response(200, request=request, json={"data": []}) + client.request = AsyncMock(return_value=success) + with patch("routstr.upstream.ppqai.time.monotonic", return_value=131.0): + assert ( + await _safe_read_request(client, "GET", str(request.url), headers={}) + is success + ) + + state = next(iter(_ppq_circuits.values())) + assert state.consecutive_failures == 0 + assert state.cooldown_until == 0.0 + + +@pytest.mark.asyncio +async def test_fetch_models_failure_is_not_a_valid_empty_catalog() -> None: + provider = PPQAIUpstreamProvider("secret") + with ( + patch( + "routstr.upstream.ppqai._safe_read_request", + AsyncMock(side_effect=httpx.ReadTimeout("catalog timed out")), + ), + pytest.raises(httpx.ReadTimeout), + ): + await provider.fetch_models() + + +@pytest.mark.asyncio +async def test_ppq_invoice_creation_post_is_never_retried() -> None: + provider = PPQAIUpstreamProvider("secret") + client = MagicMock() + client.post = AsyncMock(side_effect=httpx.ReadTimeout("invoice timed out")) + context = MagicMock() + context.__aenter__ = AsyncMock(return_value=client) + context.__aexit__ = AsyncMock(return_value=None) + + with ( + patch("routstr.upstream.ppqai.httpx.AsyncClient", return_value=context), + pytest.raises(httpx.ReadTimeout), + ): + await provider.create_lightning_topup(10, "USD") + + client.post.assert_awaited_once() diff --git a/tests/unit/test_pricing_rate_validation.py b/tests/unit/test_pricing_rate_validation.py new file mode 100644 index 00000000..96b478d1 --- /dev/null +++ b/tests/unit/test_pricing_rate_validation.py @@ -0,0 +1,489 @@ +"""Tests that an unusable rate never reaches the money math. + +A billable rate is usable only when it is finite and non-negative. Prices reach +the node from upstream catalogs, an operator's admin edit, a legacy database row +and the BTC/USD feed, and each of those can deliver ``NaN``, ``±inf`` or a +negative — ``json.loads`` accepts the bare ``NaN``/``Infinity`` literals and +overflows ``1e999`` to ``inf``. + +These tests cover the guards between such a value and a charge: the token-rate +gate that decides a model cannot be priced, the upstream-reported USD cost, the +exchange-rate feed, and the stored-row read path. They assert the node declines +to price the request rather than billing a nonsensical amount or raising after +the response has already been served. +""" + +from __future__ import annotations + +import math +from collections.abc import Iterator +from typing import Any +from unittest.mock import patch + +import pytest + +from routstr.payment.cost_calculation import ( + CostData, + MaxCostData, + calculate_cost, +) +from routstr.payment.models import ( + Architecture, + Model, + Pricing, +) + + +@pytest.fixture(autouse=True) +def patch_sats_usd_price() -> Iterator[None]: + """Pin the exchange rate; these tests are about the rates, not the feed.""" + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5): + yield + + +def _architecture() -> Architecture: + return Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ) + + +def _model(sats_pricing: Pricing) -> Model: + return Model( + id="m", + name="m", + created=0, + description="d", + context_length=8192, + architecture=_architecture(), + pricing=Pricing(prompt=1e-06, completion=2e-06), + sats_pricing=sats_pricing, + ) + + +def _usage_response() -> dict[str, Any]: + return {"model": "m", "usage": {"prompt_tokens": 1000, "completion_tokens": 500}} + + +@pytest.mark.parametrize( + "bad_rate", + [float("nan"), float("inf"), -5.0], + ids=["nan", "inf", "negative"], +) +@pytest.mark.asyncio +async def test_unusable_token_rate_never_charges_the_reservation( + bad_rate: float, +) -> None: + """An unusable configured rate must not turn authorization into usage.""" + model = _model(Pricing(prompt=bad_rate, completion=1.0)) + + cost = await calculate_cost(_usage_response(), max_cost=1234, model_obj=model) + + assert isinstance(cost, MaxCostData) + assert cost.total_msats == 0 + assert (cost.input_tokens, cost.output_tokens) == (1000, 500) + + +@pytest.mark.parametrize( + ("prompt", "completion", "expected_msats"), + [(0.0, 0.0, 0), (0.0, 2e-06, 1), (1e-06, 0.0, 1)], + ids=["free", "free-input", "free-output"], +) +@pytest.mark.asyncio +async def test_a_rate_of_zero_is_billed_as_free_not_as_missing( + prompt: float, completion: float, expected_msats: int +) -> None: + """Zero is a price, and the request must be billed on it. + + The gate that decides a model has no token pricing was a truthiness test, so + a free rate read as an absent one and the request was charged the whole + reservation instead — on a model priced at zero for that side, which is a + price the catalog serves and the router routes. + """ + model = _model(Pricing(prompt=prompt, completion=completion)) + + cost = await calculate_cost(_usage_response(), max_cost=1234, model_obj=model) + + assert isinstance(cost, CostData) + assert cost.total_msats == expected_msats + + +@pytest.mark.parametrize( + "junk", [float("inf"), float("nan"), "Infinity"], ids=["inf", "nan", "inf-string"] +) +@pytest.mark.asyncio +async def test_junk_cost_component_still_bills_the_reported_total(junk: Any) -> None: + """A malformed component must not discard the upstream's real total cost. + + ``cost_details`` only splits the total across input and output; the total is + the authoritative billed amount. A non-finite component poisoned the split, + and the request fell through to token estimation for a fraction of it. + """ + model = _model(Pricing(prompt=1e-06, completion=2e-06)) + response = { + "model": "m", + "usage": { + "prompt_tokens": 1000, + "completion_tokens": 500, + "cost": 0.01, + "cost_details": {"input_cost": junk, "output_cost": 0.004}, + }, + } + + cost = await calculate_cost(response, max_cost=9999, model_obj=model) + + assert isinstance(cost, CostData) + # $0.01 at 5.0e-5 USD/sat = 200 sats = 200_000 msats. + assert cost.total_msats == 200000 + assert cost.total_usd == pytest.approx(0.01) + + +@pytest.mark.asyncio +async def test_junk_cost_component_falls_back_to_its_alternate_field() -> None: + """A malformed component must not shadow the field that would have replaced it. + + Each side of the split has two spellings and the second is a fallback for a + missing first. ``inf`` and ``NaN`` are both truthy, so a malformed + ``input_cost`` won that choice before anything checked whether it was a + number, and the usable figure beside it was never read — the input side was + then billed at nothing and the whole total landed on output. + """ + model = _model(Pricing(prompt=1e-06, completion=2e-06)) + response = { + "model": "m", + "usage": { + "prompt_tokens": 1000, + "completion_tokens": 500, + "cost": 0.01, + "cost_details": { + "input_cost": float("inf"), + "upstream_inference_prompt_cost": 0.006, + "output_cost": 0.004, + }, + }, + } + + cost = await calculate_cost(response, max_cost=9999, model_obj=model) + + assert isinstance(cost, CostData) + # $0.01 at 5.0e-5 USD/sat = 200_000 msats, split 0.006 : 0.004. + assert cost.total_msats == 200000 + assert (cost.input_msats, cost.output_msats) == (120000, 80000) + + +@pytest.mark.asyncio +async def test_non_finite_reported_cost_is_not_a_cost() -> None: + """An upstream-reported ``Infinity`` cost is junk, not an infinite charge. + + ``json.loads`` accepts the bare ``Infinity`` literal, so a compromised or + buggy upstream can put one in ``usage.cost``. It must not be treated as a + positive USD cost at all — the request falls through to the node's own token + pricing instead. + """ + model = _model(Pricing(prompt=1e-06, completion=2e-06)) + response = { + "model": "m", + "usage": { + "prompt_tokens": 1000, + "completion_tokens": 500, + "cost": float("inf"), + }, + } + + cost = await calculate_cost(response, max_cost=9999, model_obj=model) + + assert isinstance(cost, CostData) + assert math.isfinite(cost.total_usd) + assert cost.total_msats == 2 + + +class _ExchangeResponse: + def __init__(self, payload: dict[str, Any]) -> None: + self._payload = payload + + def json(self) -> dict[str, Any]: + # A quote given as an exception stands for a response body that never + # produced one: an exchange answering with an HTML error page raises + # out of `.json()` before any price is read. + quote = next(iter(self._payload.values())) + if isinstance(quote, BaseException): + raise quote + return self._payload + + +class _ExchangeClient: + """Answers each exchange endpoint with a caller-supplied quote.""" + + def __init__(self, quotes: dict[str, Any]) -> None: + self._quotes = quotes + + async def get(self, url: str) -> _ExchangeResponse: + if "kraken" in url: + quote = self._quotes["kraken"] + if isinstance(quote, BaseException): + return _ExchangeResponse({"error": quote}) + return _ExchangeResponse({"result": {"XXBTZUSD": {"c": [quote]}}}) + if "coinbase" in url: + return _ExchangeResponse({"data": {"amount": self._quotes["coinbase"]}}) + return _ExchangeResponse({"price": self._quotes["binance"]}) + + +class _AsyncCtx: + def __init__(self, client: _ExchangeClient) -> None: + self._client = client + + async def __aenter__(self) -> _ExchangeClient: + return self._client + + async def __aexit__(self, *exc: object) -> bool: + return False + + +@pytest.fixture +def refresh_price_with() -> Iterator[Any]: + """Refresh the node's BTC/USD price from caller-supplied exchange quotes. + + Restores the module's cached price afterwards so one test cannot set the + rate another one bills at. + """ + import routstr.payment.price as price_module + + previous = (price_module.BTC_USD_PRICE, price_module.SATS_USD_PRICE) + + async def _run(quotes: dict[str, Any], last_good: float | None = None) -> None: + price_module.BTC_USD_PRICE = last_good + price_module.SATS_USD_PRICE = ( + None if last_good is None else last_good / 100_000_000 + ) + with patch.object( + price_module.httpx, + "AsyncClient", + lambda *a, **k: _AsyncCtx(_ExchangeClient(quotes)), + ): + await price_module._update_prices() + + yield _run + + price_module.BTC_USD_PRICE, price_module.SATS_USD_PRICE = previous + + +@pytest.mark.parametrize( + "bad_quote", + ["0", "0.00000000", "-1", "NaN", "Infinity", "N/A"], + ids=["zero", "zero-padded", "negative", "nan", "infinity", "non-numeric"], +) +@pytest.mark.asyncio +async def test_unusable_exchange_quote_does_not_set_the_node_price( + bad_quote: str, refresh_price_with: Any +) -> None: + """One exchange returning junk must not set the price the node bills at. + + The feed takes the ``min()`` of what it collects, so an unusable quote does + not merely join the sample — it *wins*. The two healthy quotes must still + price the node. + """ + from routstr.payment.price import btc_usd_price + + await refresh_price_with( + {"kraken": bad_quote, "coinbase": "100000.0", "binance": "100000.0"} + ) + + assert btc_usd_price() == pytest.approx(100000.0) + + +@pytest.mark.asyncio +async def test_boolean_exchange_quote_does_not_set_the_node_price( + refresh_price_with: Any, +) -> None: + """A boolean in the price field is a shape change, not a $1 bitcoin. + + ``float(True)`` is ``1.0``, which is finite and positive, so a payload whose + price field turned into a boolean passes every numeric guard — and then + *wins* the ``min()``, pricing the whole node at one dollar per bitcoin. + """ + from routstr.payment.price import btc_usd_price + + await refresh_price_with( + {"kraken": True, "coinbase": "100000.0", "binance": "100000.0"} + ) + + assert btc_usd_price() == pytest.approx(100000.0) + + +@pytest.mark.asyncio +async def test_an_underflowing_exchange_quote_does_not_set_the_node_price( + refresh_price_with: Any, +) -> None: + """A quote too small to survive the sats conversion is not a price. + + ``1e-320`` is positive, so it passes the guards and wins the ``min()``, but + the node prices in sats and ``1e-320 / 100_000_000`` underflows to ``0.0`` + — a zero sats price divides by zero on every model's rate. + """ + from routstr.payment.price import btc_usd_price + + await refresh_price_with( + {"kraken": "1e-320", "coinbase": "100000.0", "binance": "100000.0"} + ) + + assert btc_usd_price() == pytest.approx(100000.0) + + +@pytest.mark.asyncio +async def test_an_unreadable_exchange_response_drops_only_that_quote( + refresh_price_with: Any, +) -> None: + """An exchange whose response never yields a quote costs one quote. + + The price is aggregated across three exchanges so that one of them having a + bad day is survivable; unhandled, the raise aborted the whole aggregation. + """ + from routstr.payment.price import btc_usd_price + + await refresh_price_with( + { + "kraken": ValueError("Expecting value: line 1 column 1 (char 0)"), + "coinbase": "100000.0", + "binance": "100000.0", + } + ) + + assert btc_usd_price() == pytest.approx(100000.0) + + +@pytest.mark.asyncio +async def test_all_quotes_unusable_keeps_the_last_good_price( + refresh_price_with: Any, +) -> None: + """When every quote is junk the node keeps the last price it trusted. + + Adopting ``0`` or ``NaN`` because it was the only thing on offer would take + out billing for every model at once; skipping the update degrades to a stale + rate, which is the safe direction and what an unreachable exchange already + does. + """ + from routstr.payment.price import btc_usd_price + + await refresh_price_with( + {"kraken": "0", "coinbase": "NaN", "binance": "-3"}, last_good=90000.0 + ) + + assert btc_usd_price() == pytest.approx(90000.0) + + +# --------------------------------------------------------------------------- +# Catalog ingest — a malformed rate must never become a stored price +# --------------------------------------------------------------------------- + + +class _CatalogResponse: + 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 _CatalogClient: + """Stands in for ``httpx.AsyncClient`` against the OpenRouter catalog.""" + + def __init__(self, models: list[dict[str, Any]]) -> None: + self._models = models + + async def __aenter__(self) -> "_CatalogClient": + return self + + async def __aexit__(self, *exc: object) -> bool: + return False + + async def get(self, url: str, timeout: int | None = None) -> _CatalogResponse: + if url.endswith("/embeddings/models"): + return _CatalogResponse({"data": []}) + return _CatalogResponse({"data": self._models}) + + +def _catalog_entry(model_id: str, pricing: dict[str, Any]) -> dict[str, Any]: + return {"id": model_id, "name": model_id, "pricing": pricing} + + +def _patch_openrouter_catalog(models: list[dict[str, Any]]) -> Any: + return patch( + "routstr.payment.models.httpx.AsyncClient", + lambda *args, **kwargs: _CatalogClient(models), + ) + + +@pytest.mark.parametrize( + "bad_rate", + [float("nan"), float("inf"), float("-inf")], + ids=["nan", "inf", "negative-inf"], +) +@pytest.mark.asyncio +async def test_non_finite_catalog_rate_is_not_imported(bad_rate: float) -> None: + """A non-finite rate in the upstream catalog is junk, not a price. + + ``json.loads`` accepts the bare ``NaN``/``Infinity`` literals and overflows + ``1e999`` to ``inf``, so an upstream feed can deliver one. The import filter + rejects a negative and a both-zero price, but every comparison with ``NaN`` + is False and ``inf`` reads as a large positive, so both sailed through and + became a stored price the node would advertise and bill on. + """ + with _patch_openrouter_catalog( + [ + _catalog_entry("bad", {"prompt": bad_rate, "completion": "0.000002"}), + _catalog_entry("good", {"prompt": "0.000001", "completion": "0.000002"}), + ] + ): + from routstr.payment.models import async_fetch_openrouter_models + + models = await async_fetch_openrouter_models() + + assert [m["id"] for m in models] == ["good"] + + +@pytest.mark.asyncio +async def test_oversized_catalog_rate_does_not_empty_the_catalog() -> None: + """An integer too large to be a float must cost one model, not all of them. + + ``float()`` raises ``OverflowError`` — not ``ValueError`` — for such a + value, so the coercion guard in the import filter did not catch it and the + exception unwound the whole fetch. The node then imported nothing at all + from an upstream whose catalog was fine apart from one entry. + """ + with _patch_openrouter_catalog( + [ + _catalog_entry("bad", {"prompt": 10**400, "completion": 2}), + _catalog_entry("good", {"prompt": "0.000001", "completion": "0.000002"}), + ] + ): + from routstr.payment.models import async_fetch_openrouter_models + + models = await async_fetch_openrouter_models() + + assert [m["id"] for m in models] == ["good"] + + +@pytest.mark.asyncio +async def test_boolean_catalog_rate_is_not_imported() -> None: + """A JSON ``true`` is a change of shape, not a price. + + Python coerces it to a finite, positive ``1.0`` — a dollar per token — so it + passes every numeric guard and must be rejected before coercion. + """ + with _patch_openrouter_catalog( + [ + _catalog_entry("bad", {"prompt": True, "completion": "0.000002"}), + _catalog_entry("good", {"prompt": "0.000001", "completion": "0.000002"}), + ] + ): + from routstr.payment.models import async_fetch_openrouter_models + + models = await async_fetch_openrouter_models() + + assert [m["id"] for m in models] == ["good"] diff --git a/tests/unit/test_proxy_path_allowlist.py b/tests/unit/test_proxy_path_allowlist.py new file mode 100644 index 00000000..4cd1c67f --- /dev/null +++ b/tests/unit/test_proxy_path_allowlist.py @@ -0,0 +1,217 @@ +"""Unit tests for the proxy edge path allowlist (arbitrary-upstream-path-proxy). + +An authenticated POST used to be forwarded for ANY path, so a caller could reach +arbitrary or traversal-shaped upstream endpoints with the provider credential +attached. The proxy now rejects ambiguous path spellings for every method, then +requires the method/path pair to name a canonical endpoint. A familiar prefix is +no longer enough: "v1/organization/api_keys" is refused just like "internal/admin". +""" + +from __future__ import annotations + +import os + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +import pytest # noqa: E402 + +from routstr.proxy import ( # noqa: E402 + _forwarding_allowed, + _is_ambiguously_spelled_path, + _parse_extra_allowed_endpoints, +) + + +@pytest.mark.parametrize( + "path", + [ + "../secret", + "v1/../admin", + "v1/./models", + "..", + "v1//models", # duplicate separator + "/v1/models", # leading slash / absolute override + "v1/models/..", + "%2e%2e/secret", # residual encoded dot segment + "v1/%2fadmin", # residual encoded slash + "v1\\models", # backslash + "v1/models\x00", # NUL byte + " v1/models", # leading whitespace + "", + ], +) +def test_ambiguous_paths_are_rejected(path: str) -> None: + assert _is_ambiguously_spelled_path(path) is True + + +@pytest.mark.parametrize( + "path", + [ + "v1/chat/completions", + "chat/completions", + "v1/responses", + "v1/embeddings", + "models", + "v1/models/gpt-4", + "attestation/", # a single trailing slash is canonical + "tee/attestation/", + ], +) +def test_canonical_paths_are_allowed(path: str) -> None: + assert _is_ambiguously_spelled_path(path) is False + + +def test_unknown_paths_are_not_forwarded() -> None: + # The credential is attached during forwarding, so an unknown endpoint must + # never be forwarded on the caller's say-so. + assert _forwarding_allowed("internal/admin", "POST") is False + assert _forwarding_allowed("secret-endpoint", "POST") is False + + +@pytest.mark.parametrize( + "path", + [ + "moderations", + "rerank", + "audio/speech", + "audio/transcriptions", + "audio/translations", + "images/generations", + "images/edits", + "images/variations", + ], +) +def test_unbilled_endpoints_are_not_forwarded_by_default(path: str) -> None: + assert _forwarding_allowed(path, "POST") is False + assert _forwarding_allowed(f"v1/{path}", "POST") is False + + +@pytest.mark.parametrize( + "path", + [ + "modelsdump", # "models" must not match a longer segment + "attestationadmin", + "providers-secret", + "embeddingsx", + "completions-internal", + ], +) +def test_endpoint_name_does_not_match_a_longer_segment(path: str) -> None: + assert _forwarding_allowed(path, "POST") is False + assert _forwarding_allowed(path, "GET") is False + + +@pytest.mark.parametrize( + "path", + [ + # A familiar prefix must not carry an unknown endpoint. These are real + # upstream routes that manage keys, org membership, and billing. + "v1/organization/api_keys", + "v1/api_keys", + "v1/admin/keys", + "v1/billing/usage", + "v1/files", + "v1/batches", + "chat/internal", + "audio/internal", + "images/internal", + "tee/keys", + # No endpoint takes a trailing id segment; a resource id never widens + # the reachable surface. + "models/gpt-4", + "models/gpt-4/secret", + "chat/completions/abc", + ], +) +def test_known_prefix_does_not_carry_an_unknown_endpoint(path: str) -> None: + assert _forwarding_allowed(path, "POST") is False + assert _forwarding_allowed(path, "GET") is False + + +@pytest.mark.parametrize( + ("path", "method"), + [ + ("v1/chat/completions", "POST"), + ("chat/completions", "POST"), + ("v1/chat/completions/", "POST"), + ("completions", "POST"), + ("v1/responses", "POST"), + ("v1/messages", "POST"), + ("v1/embeddings", "POST"), + ("models", "GET"), + ("attestation", "GET"), + ("tee/attestation", "GET"), + ], +) +def test_canonical_endpoints_are_forwarded(path: str, method: str) -> None: + assert _forwarding_allowed(path, method) is True + + +@pytest.mark.parametrize( + ("path", "method"), + [ + ("chat/completions", "GET"), # billed endpoints are POST-only + ("v1/embeddings", "GET"), + ("models", "POST"), # read-only endpoints are GET-only + ("attestation", "POST"), + ("v1/chat/completions", "DELETE"), # never routed here, refused anyway + ("v1/chat/completions", "PUT"), + ], +) +def test_method_must_match_the_endpoint(path: str, method: str) -> None: + assert _forwarding_allowed(path, method) 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 == { + "rerank": frozenset({"POST"}), # the "v1/" prefix collapses like any path + "batches": frozenset({"GET"}), + "audio/x": frozenset({"POST"}), + } + + +def test_operator_additions_may_grant_two_methods_on_one_endpoint() -> None: + assert _parse_extra_allowed_endpoints("POST:batches,GET:batches") == { + "batches": frozenset({"POST", "GET"}) + } + + +@pytest.mark.parametrize( + "raw", + [ + "", + " ", + "v1/rerank", # no method + "POST:", # no path + ":v1/rerank", # empty method + "DELETE:v1/rerank", # method the proxy never routes + "POST:*", # wildcards are deliberately unsupported + "POST:v1/*", + "POST:../secret", # ambiguous spellings are screened here too + "POST:v1//rerank", + "POST:%2e%2e/secret", + ], +) +def test_malformed_operator_additions_widen_nothing(raw: str) -> None: + assert _parse_extra_allowed_endpoints(raw) == {} + + +def test_operator_additions_are_env_only() -> None: + # The proxy parses this once at import, so a persisted or admin-API-writable + # value would be read but never take effect. Keeping it env-only also means + # widening the reachable upstream surface takes a deploy. + from routstr.core.settings import ENV_ONLY_FIELDS + + assert "proxy_extra_allowed_paths" in ENV_ONLY_FIELDS + + +def test_ehbp_is_gated_by_the_same_allowlist() -> None: + # EHBP hides the request body from the proxy, which is a reason to constrain + # the destination more tightly rather than to trust the caller's path: the + # encrypted contract covers the body, never the endpoint the provider + # credential is spent against. + assert _forwarding_allowed("anything/encrypted", "POST") is False + assert _forwarding_allowed("v1/organization/api_keys", "POST") is False + assert _forwarding_allowed("v1/chat/completions", "POST") is True diff --git a/tests/unit/test_proxy_tinfoil_attestation_routing.py b/tests/unit/test_proxy_tinfoil_attestation_routing.py index b6ba2f88..c367a049 100644 --- a/tests/unit/test_proxy_tinfoil_attestation_routing.py +++ b/tests/unit/test_proxy_tinfoil_attestation_routing.py @@ -102,10 +102,22 @@ async def test_attestation_trailing_slash_routes_directly_to_tinfoil( tinfoil.forward_get_request.assert_awaited_once() -@pytest.mark.parametrize("path", ["attestation/foo", "attestationjunk"]) +@pytest.mark.parametrize( + "path", + [ + # A valid `attestation` segment is not the exact attestation route, and + # `attestation` takes no id segment, so the endpoint allowlist rejects + # it at the edge rather than letting it reach model/auth handling. + "attestation/foo", + # Not a known endpoint at all: rejected at the edge before routing. + "attestationjunk", + ], +) @pytest.mark.asyncio async def test_non_attestation_prefix_does_not_bypass_authentication( - monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI, path: str + monkeypatch: pytest.MonkeyPatch, + proxy_app: FastAPI, + path: str, ) -> None: tinfoil = MagicMock() tinfoil.provider_type = "tinfoil" @@ -118,8 +130,7 @@ async def test_non_attestation_prefix_does_not_bypass_authentication( ) as client: response = await client.get(f"/{path}") - assert response.status_code == 400 - assert response.json()["error"]["type"] == "invalid_model" + assert response.status_code == 404 tinfoil.forward_get_request.assert_not_awaited() diff --git a/tests/unit/test_ranking_maxcost_mismatch.py b/tests/unit/test_ranking_maxcost_mismatch.py new file mode 100644 index 00000000..438dcf0d --- /dev/null +++ b/tests/unit/test_ranking_maxcost_mismatch.py @@ -0,0 +1,169 @@ +"""Ranking-vs-reservation mismatch for same-model multi-provider setups. + +``calculate_model_cost_score`` weights a typical request and ignores +``context_length``, while the balance gate reserves on the context-based +``sats_pricing.max_cost``. When two providers serve the same model these can +disagree, so the "cheapest" advertised provider may demand a far larger +reservation and surprise the client with a 402. Ranking must therefore use the +same context-based ceiling as the gate. +""" + +import os +from unittest.mock import Mock + +import pytest + +os.environ["UPSTREAM_BASE_URL"] = "http://test" +os.environ["UPSTREAM_API_KEY"] = "test" + +from routstr.algorithm import ( # noqa: E402 + calculate_model_cost_score, + create_model_mappings, +) +from routstr.payment.helpers import get_max_cost_for_model # noqa: E402 +from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 + + +def _arch() -> Architecture: + return Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="gpt", + instruct_type=None, + ) + + +def _model( + model_id: str, + prompt: float, + completion: float, + context_length: int, + max_cost_sats: float, +) -> Model: + """Build a Model whose sats_pricing.max_cost mirrors the context-based gate.""" + pricing = Pricing( + prompt=prompt, + completion=completion, + request=0.0, + image=0.0, + web_search=0.0, + internal_reasoning=0.0, + ) + model = Model( + id=model_id, + name=model_id, + created=1, + description="", + context_length=context_length, + architecture=_arch(), + pricing=pricing, + ) + model.sats_pricing = Pricing( + prompt=prompt, + completion=completion, + request=0.0, + image=0.0, + web_search=0.0, + internal_reasoning=0.0, + max_prompt_cost=context_length * prompt, + max_completion_cost=context_length * completion, + max_cost=max_cost_sats, + ) + return model + + +def _provider(name: str, db_id: int, models: list[Model]) -> Mock: + provider = Mock() + provider.provider_type = name + provider.base_url = f"https://{name}.example/v1" + provider.db_id = db_id + provider.upstream_name = name + provider.provider_fee = 1.0 + provider.get_cached_models.return_value = models + return provider + + +# Provider LOW-SCORE: cheaper per typical token, but a huge context window makes +# its context-based max_cost enormous. +_LOW_SCORE = _model( + "shared-model", + prompt=0.001, + completion=0.001, + context_length=1_000_000, + max_cost_sats=1000.0, +) +# Provider LOW-RESERVE: pricier per typical token, but a small context window +# makes its reservation ceiling tiny. +_LOW_RESERVE = _model( + "shared-model", + prompt=0.002, + completion=0.002, + context_length=8_000, + max_cost_sats=16.0, +) + + +def test_ranking_metric_and_reservation_metric_disagree() -> None: + """The two cost metrics rank the same two providers in opposite order.""" + assert calculate_model_cost_score(_LOW_SCORE) < calculate_model_cost_score( + _LOW_RESERVE + ) + assert _max_cost(_LOW_SCORE) > _max_cost(_LOW_RESERVE) + + +@pytest.mark.asyncio +async def test_get_max_cost_uses_context_based_max_cost_not_score() -> None: + """The balance gate reserves on max_cost, so the low-score model costs more.""" + session = Mock() + low_score_reserve = await get_max_cost_for_model( + "shared-model", session=session, model_obj=_LOW_SCORE + ) + low_reserve_reserve = await get_max_cost_for_model( + "shared-model", session=session, model_obj=_LOW_RESERVE + ) + # The "cheapest" model (by ranking score) demands the LARGER reservation. + assert low_score_reserve > low_reserve_reserve + + +def _max_cost(model: Model) -> float: + assert model.sats_pricing is not None + assert model.sats_pricing.max_cost is not None + return model.sats_pricing.max_cost + + +def _both_orderings() -> list[list[BaseUpstreamProvider]]: + low_score = _provider("low-score", 1, [_LOW_SCORE]) + low_reserve = _provider("low-reserve", 2, [_LOW_RESERVE]) + return [[low_score, low_reserve], [low_reserve, low_score]] + + +def test_catalog_advertises_the_lower_reservation_provider() -> None: + """Catalog and routing pick the provider with the smaller reservation, + regardless of upstream iteration order.""" + for providers in _both_orderings(): + _, provider_map, unique_models = create_model_mappings( + upstreams=providers, + overrides_by_key={}, + disabled_model_keys=set(), + ) + selected_model, selected_provider = provider_map["shared-model"][0] + assert selected_provider.provider_type == "low-reserve" + assert _max_cost(selected_model) == 16.0 + assert _max_cost(unique_models["shared-model"]) == 16.0 + + +def test_selected_candidate_has_minimal_reservation() -> None: + """The first-tried candidate never demands a larger reservation than + another available candidate for the same model.""" + for providers in _both_orderings(): + _, provider_map, _ = create_model_mappings( + upstreams=providers, + overrides_by_key={}, + disabled_model_keys=set(), + ) + candidates = provider_map["shared-model"] + selected_max_cost = _max_cost(candidates[0][0]) + min_max_cost = min(_max_cost(m) for m, _ in candidates) + assert selected_max_cost == min_max_cost diff --git a/tests/unit/test_reasoning_effort.py b/tests/unit/test_reasoning_effort.py new file mode 100644 index 00000000..55d3b980 --- /dev/null +++ b/tests/unit/test_reasoning_effort.py @@ -0,0 +1,264 @@ +"""Per-model reasoning-effort catalog metadata and request mapping.""" + +from __future__ import annotations + +import json +import os +from typing import Any + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") +os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to") + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from routstr.core.db import get_session +from routstr.payment.models import ( + Architecture, + Model, + Pricing, + Reasoning, + models_router, +) +from routstr.upstream import GenericUpstreamProvider +from routstr.upstream.reasoning_effort import ( + adapt_messages_body_for_litellm, + apply_reasoning_effort, + closest_supported_effort, + extract_requested_effort, + resolve_effort, +) + + +def _model(**kwargs: Any) -> Model: + reasoning = kwargs.pop("reasoning", None) + return Model( + id=kwargs.get("id", "openai/gpt-5.6-sol"), + 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), + reasoning=reasoning, + ) + + +SOL_REASONING = { + "mandatory": False, + "default_enabled": True, + "supported_efforts": ["max", "xhigh", "high", "medium", "low", "none"], + "default_effort": "medium", +} + + +def test_model_parses_openrouter_reasoning_object() -> None: + model = Model.parse_obj( + { + "id": "openai/gpt-5.6-sol", + "name": "GPT", + "created": 0, + "description": "", + "context_length": 1, + "architecture": { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "x", + "instruct_type": None, + }, + "pricing": {"prompt": 1e-6, "completion": 1e-6}, + "reasoning": SOL_REASONING, + "extra_ignored_field": "drop me", + } + ) + assert model.reasoning is not None + assert model.reasoning.supported_efforts == [ + "max", + "xhigh", + "high", + "medium", + "low", + "none", + ] + dumped = model.dict() + assert dumped["reasoning"]["supported_efforts"][0] == "max" + assert dumped["reasoning"]["default_effort"] == "medium" + assert "extra_ignored_field" not in dumped + + +def test_non_reasoning_models_omit_the_field() -> None: + dumped = _model().dict() + assert "reasoning" not in dumped + + +def test_malformed_reasoning_is_dropped_not_fatal() -> None: + model = _model(reasoning=["not", "a", "dict"]) + assert model.reasoning is None + assert "reasoning" not in model.dict() + + +def test_closest_effort_maps_minimal_to_low() -> None: + assert ( + closest_supported_effort( + "minimal", + ["max", "xhigh", "high", "medium", "low", "none"], + default_effort="medium", + ) + == "low" + ) + + +def test_closest_effort_maps_max_when_missing() -> None: + assert ( + closest_supported_effort( + "max", + ["high", "medium", "low", "none"], + default_effort="medium", + ) + == "high" + ) + + +def test_mandatory_rejects_none() -> None: + assert ( + closest_supported_effort( + "none", + ["max", "high", "medium", "low", "none"], + default_effort="high", + mandatory=True, + ) + == "high" + ) + + +def test_missing_request_uses_default() -> None: + reasoning = Reasoning.parse_obj(SOL_REASONING) + assert resolve_effort(None, reasoning) == "medium" + + +def test_extract_prefers_nested_reasoning_effort() -> None: + assert ( + extract_requested_effort( + {"reasoning_effort": "low", "reasoning": {"effort": "high"}} + ) + == "high" + ) + + +def test_prepare_request_body_rewrites_unsupported_effort() -> None: + provider = GenericUpstreamProvider(base_url="https://openrouter.ai/api/v1") + model = _model(reasoning=SOL_REASONING) + body = json.dumps( + { + "model": "openai/gpt-5.6-sol", + "messages": [{"role": "user", "content": "hi"}], + "reasoning_effort": "minimal", + } + ).encode() + out = provider.prepare_request_body(body, model) + assert out is not None + data = json.loads(out) + assert data["reasoning_effort"] == "low" + + +def test_prepare_request_body_rewrites_nested_effort() -> None: + provider = GenericUpstreamProvider(base_url="https://openrouter.ai/api/v1") + model = _model(reasoning=SOL_REASONING) + body = json.dumps( + { + "model": "openai/gpt-5.6-sol", + "messages": [{"role": "user", "content": "hi"}], + "reasoning": {"effort": "minimal", "exclude": False}, + } + ).encode() + out = provider.prepare_request_body(body, model) + assert out is not None + data = json.loads(out) + assert data["reasoning"]["effort"] == "low" + assert data["reasoning"]["exclude"] is False + + +def test_prepare_request_body_leaves_plain_chat_alone() -> None: + provider = GenericUpstreamProvider(base_url="https://openrouter.ai/api/v1") + model = _model(reasoning=SOL_REASONING) + payload = { + "model": "openai/gpt-5.6-sol", + "messages": [{"role": "user", "content": "hi"}], + } + body = json.dumps(payload).encode() + out = provider.prepare_request_body(body, model) + assert out == body + + +def test_apply_injects_default_when_mandatory() -> None: + data: dict[str, Any] = { + "model": "anthropic/claude-fable-5.1", + "messages": [{"role": "user", "content": "hi"}], + } + model = _model( + id="anthropic/claude-fable-5.1", + reasoning={ + "mandatory": True, + "supported_efforts": ["max", "xhigh", "high", "medium", "low"], + "default_effort": "high", + }, + ) + assert apply_reasoning_effort(data, model) is True + assert data["reasoning_effort"] == "high" + + +def test_fee_apply_preserves_reasoning() -> None: + provider = GenericUpstreamProvider( + base_url="https://openrouter.ai/api/v1", provider_fee=1.1 + ) + model = _model(reasoning=SOL_REASONING) + priced = provider._apply_provider_fee_to_model(model) + assert priced.reasoning is not None + assert priced.reasoning.supported_efforts == SOL_REASONING["supported_efforts"] + + +def test_v1_models_includes_reasoning_and_omits_when_absent( + monkeypatch: Any, +) -> None: + with_reasoning = _model(id="openai/gpt-5.6-sol", reasoning=SOL_REASONING) + without = _model(id="openai/gpt-4o") + unique = {"openai/gpt-5.6-sol": with_reasoning, "openai/gpt-4o": without} + + import routstr.proxy as proxy + + monkeypatch.setattr(proxy, "_unique_models", unique) + app = FastAPI() + app.include_router(models_router) + app.dependency_overrides[get_session] = lambda: None + response = TestClient(app).get("/v1/models") + assert response.status_code == 200 + by_id = {row["id"]: row for row in response.json()["data"]} + assert by_id["openai/gpt-5.6-sol"]["reasoning"]["supported_efforts"] == [ + "max", + "xhigh", + "high", + "medium", + "low", + "none", + ] + assert "reasoning" not in by_id["openai/gpt-4o"] + + +def test_messages_thinking_becomes_reasoning_effort() -> None: + model = _model(reasoning=SOL_REASONING) + body: dict[str, Any] = { + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 16, + "thinking": {"type": "enabled", "effort": "minimal"}, + } + adapt_messages_body_for_litellm(body, model) + assert "thinking" not in body + assert body["reasoning_effort"] == "low" diff --git a/tests/unit/test_redemption_negative_cache.py b/tests/unit/test_redemption_negative_cache.py new file mode 100644 index 00000000..5d19d411 --- /dev/null +++ b/tests/unit/test_redemption_negative_cache.py @@ -0,0 +1,148 @@ +"""Unit tests for the terminal-redemption negative cache.""" + +from typing import Iterator + +import pytest +from fastapi import HTTPException + +from routstr.auth import ( + _cached_failure_to_http_exception, + _maybe_cache_terminal_redemption_failure, +) +from routstr.redemption_cache import ( + TERMINAL_REDEMPTION_CODES, + CachedRedemptionFailure, + RedemptionNegativeCache, + redemption_negative_cache, +) + +FAILURE = CachedRedemptionFailure( + status_code=400, + error_type="token_already_spent", + message="Cashu token already spent", + code="cashu_token_already_spent", +) + + +class FakeClock: + def __init__(self) -> None: + self.now = 0.0 + + def __call__(self) -> float: + return self.now + + +@pytest.fixture(autouse=True) +def _clean_singleton() -> Iterator[None]: + redemption_negative_cache.clear() + yield + redemption_negative_cache.clear() + + +class TestRedemptionNegativeCache: + def test_get_returns_none_for_unknown_key(self) -> None: + cache = RedemptionNegativeCache() + assert cache.get("deadbeef") is None + + def test_put_then_get_roundtrip(self) -> None: + cache = RedemptionNegativeCache() + cache.put("deadbeef", FAILURE) + assert cache.get("deadbeef") == FAILURE + + def test_entry_expires_after_ttl(self) -> None: + clock = FakeClock() + cache = RedemptionNegativeCache(ttl_seconds=100, clock=clock) + cache.put("deadbeef", FAILURE) + clock.now = 99.9 + assert cache.get("deadbeef") == FAILURE + clock.now = 100.0 + assert cache.get("deadbeef") is None + assert len(cache) == 0 + + def test_lru_eviction_at_capacity(self) -> None: + cache = RedemptionNegativeCache(max_entries=2) + cache.put("a", FAILURE) + cache.put("b", FAILURE) + # Touch "a" so "b" becomes the least recently used entry. + assert cache.get("a") is not None + cache.put("c", FAILURE) + assert cache.get("b") is None + assert cache.get("a") is not None + assert cache.get("c") is not None + + def test_reput_refreshes_expiry(self) -> None: + clock = FakeClock() + cache = RedemptionNegativeCache(ttl_seconds=100, clock=clock) + cache.put("deadbeef", FAILURE) + clock.now = 90.0 + cache.put("deadbeef", FAILURE) + clock.now = 150.0 + assert cache.get("deadbeef") == FAILURE + + def test_discard_removes_entry(self) -> None: + cache = RedemptionNegativeCache() + cache.put("deadbeef", FAILURE) + cache.discard("deadbeef") + assert cache.get("deadbeef") is None + cache.discard("deadbeef") # idempotent + + def test_invalid_construction_args_rejected(self) -> None: + with pytest.raises(ValueError): + RedemptionNegativeCache(max_entries=0) + with pytest.raises(ValueError): + RedemptionNegativeCache(ttl_seconds=0) + + +class TestMaybeCacheTerminalRedemptionFailure: + def test_already_spent_error_is_cached(self) -> None: + _maybe_cache_terminal_redemption_failure( + "deadbeef", Exception("Mint Error: Token already spent. (Code: 11001)") + ) + cached = redemption_negative_cache.get("deadbeef") + assert cached is not None + assert cached.code == "cashu_token_already_spent" + assert cached.status_code == 400 + + def test_transient_mint_unreachable_is_not_cached(self) -> None: + import httpx + + _maybe_cache_terminal_redemption_failure( + "deadbeef", httpx.ConnectError("connection refused") + ) + assert redemption_negative_cache.get("deadbeef") is None + + def test_unclassified_error_is_not_cached(self) -> None: + _maybe_cache_terminal_redemption_failure( + "deadbeef", RuntimeError("some internal fault") + ) + assert redemption_negative_cache.get("deadbeef") is None + + def test_generic_value_error_is_not_cached(self) -> None: + # cashu_token_redemption_failed is deliberately NOT terminal — a + # generic ValueError can wrap transient faults. + _maybe_cache_terminal_redemption_failure( + "deadbeef", ValueError("something went wrong during redemption") + ) + assert redemption_negative_cache.get("deadbeef") is None + + def test_terminal_codes_are_a_closed_set(self) -> None: + assert TERMINAL_REDEMPTION_CODES == { + "cashu_token_already_spent", + "invalid_cashu_token", + "cashu_token_zero_value", + "cashu_token_swap_fees_exceed_amount", + } + + +class TestCachedFailureToHttpException: + def test_envelope_matches_classify_taxonomy(self) -> None: + exc = _cached_failure_to_http_exception(FAILURE) + assert isinstance(exc, HTTPException) + assert exc.status_code == 400 + assert exc.detail == { + "error": { + "message": "Cashu token already spent", + "type": "token_already_spent", + "code": "cashu_token_already_spent", + } + } diff --git a/tests/unit/test_refund_no_retry.py b/tests/unit/test_refund_no_retry.py new file mode 100644 index 00000000..26ab0916 --- /dev/null +++ b/tests/unit/test_refund_no_retry.py @@ -0,0 +1,47 @@ +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +from fastapi import HTTPException + +from routstr.upstream.base import BaseUpstreamProvider + + +@pytest.mark.asyncio +async def test_send_refund_does_not_retry_ambiguous_token_creation() -> None: + provider = object.__new__(BaseUpstreamProvider) + send_token = AsyncMock(side_effect=httpx.ReadTimeout("swap response lost")) + store = AsyncMock() + + with ( + patch("routstr.upstream.base.send_token", send_token), + patch("routstr.upstream.base.store_cashu_transaction", store), + pytest.raises(HTTPException) as raised, + ): + await provider.send_refund(10, "sat", mint="https://mint.test") + + assert raised.value.status_code == 401 + send_token.assert_awaited_once_with( + 10, unit="sat", mint_url="https://mint.test" + ) + store.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_send_refund_does_not_retry_or_store_on_generic_failure() -> None: + provider = object.__new__(BaseUpstreamProvider) + send_token = AsyncMock(side_effect=Exception("mint rejected swap")) + store = AsyncMock() + + with ( + patch("routstr.upstream.base.send_token", send_token), + patch("routstr.upstream.base.store_cashu_transaction", store), + pytest.raises(HTTPException) as raised, + ): + await provider.send_refund(10, "sat", mint="https://mint.test") + + assert raised.value.status_code == 401 + send_token.assert_awaited_once_with( + 10, unit="sat", mint_url="https://mint.test" + ) + store.assert_not_awaited() diff --git a/tests/unit/test_refund_script_url_policy.py b/tests/unit/test_refund_script_url_policy.py new file mode 100644 index 00000000..2ca8d502 --- /dev/null +++ b/tests/unit/test_refund_script_url_policy.py @@ -0,0 +1,49 @@ +"""The refund helper puts a cashu token and a bearer key on the wire, so it +must not speak cleartext to a remote host.""" + +import importlib.util +import sys +from pathlib import Path +from types import ModuleType + +import pytest + +SCRIPT = ( + Path(__file__).resolve().parents[2] / "scripts" / "refund_token_to_lightning.py" +) + + +def _load() -> ModuleType: + spec = importlib.util.spec_from_file_location("refund_token_to_lightning", SCRIPT) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +@pytest.mark.parametrize( + "url", + [ + "https://node.example.com", + "http://localhost:8000", + "http://127.0.0.1:8000", + "http://[::1]:8000", + ], +) +def test_accepts_https_and_loopback_http(url: str) -> None: + assert _load().check_url(url) == url + + +@pytest.mark.parametrize( + "url", + [ + "http://node.example.com", + "http://192.168.1.10:8000", + "ftp://node.example.com", + "node.example.com", + ], +) +def test_rejects_remote_cleartext_and_other_schemes(url: str) -> None: + with pytest.raises(SystemExit): + _load().check_url(url) diff --git a/tests/unit/test_request_correction.py b/tests/unit/test_request_correction.py index 903fe16e..3b1110a2 100644 --- a/tests/unit/test_request_correction.py +++ b/tests/unit/test_request_correction.py @@ -63,7 +63,9 @@ class TestCorrectRequest: assert correct_request(_body(temperature=1), "", set()) is None def test_returns_none_on_non_object_body(self) -> None: - assert correct_request(b"[1, 2, 3]", "`temperature` is deprecated", set()) is None + assert ( + correct_request(b"[1, 2, 3]", "`temperature` is deprecated", set()) is None + ) def test_deprecated_model_name_is_not_stripped_as_param(self) -> None: """A 'model is deprecated' error must not strip an unrelated body field. @@ -106,6 +108,36 @@ class TestStripUnsupportedParam: def test_declines_when_no_match(self) -> None: assert strip_unsupported_param({"temperature": 1}, "nope") is None + def test_never_strips_spend_shaping_params(self) -> None: + # Stripping an output cap after the reservation was priced would let + # the retry run uncapped and overcharge — the corrector must decline so + # the upstream error propagates instead. + for param in ( + "max_tokens", + "max_completion_tokens", + "max_output_tokens", + "max_tokens_to_sample", + "n", + "best_of", + ): + body = {"model": "m", param: 4, "messages": []} + assert ( + strip_unsupported_param(body, f"`{param}` is not supported") is None + ), param + # And through the full pipeline entry point. + assert ( + correct_request( + json.dumps(body).encode(), + f"`{param}` is not supported", + set(), + ) + is None + ), param + + def test_spend_shaping_guard_is_case_insensitive(self) -> None: + body = {"model": "m", "Max_Tokens": 4} + assert strip_unsupported_param(body, "`Max_Tokens` is deprecated") is None + class TestExtractErrorMessage: def test_extracts_nested_error_message(self) -> None: diff --git a/tests/unit/test_short_keyset_ids.py b/tests/unit/test_short_keyset_ids.py index 870e76d3..7f87b4de 100644 --- a/tests/unit/test_short_keyset_ids.py +++ b/tests/unit/test_short_keyset_ids.py @@ -1,11 +1,8 @@ -from typing import cast -from unittest.mock import AsyncMock, Mock, patch +from unittest.mock import AsyncMock, Mock import httpx import pytest from cashu.core.base import ( - MeltQuoteState, - Proof, TokenV4, TokenV4Proof, TokenV4Token, @@ -16,7 +13,6 @@ from routstr.wallet import ( Wallet, _redeem_same_mint, classify_redemption_error, - swap_to_trusted_mint, ) MINT_URL = "https://mint.example" @@ -163,54 +159,3 @@ async def test_cached_keysets_do_not_mask_a_refresh_failure( assert classified is not None assert classified[1] == 503 assert classified[3] == "cashu_source_mint_unreachable" - - -@pytest.mark.asyncio -async def test_cross_mint_swap_uses_resolved_proofs_and_active_output_keyset() -> None: - token = _token(amounts=(7,)) - source_wallet = _wallet_with_keysets(FULL_V2_ID) - source_wallet.melt_quote = AsyncMock( - return_value=Mock(quote="melt-quote", amount=5, fee_reserve=2) - ) - - async def assert_melt_boundary(**kwargs: object) -> Mock: - assert kwargs["fee_reserve_sat"] == 2 - assert source_wallet.keyset_id == FULL_V2_ID - return Mock(state=MeltQuoteState.paid) - - source_wallet.melt = AsyncMock(side_effect=assert_melt_boundary) - - destination_url = "https://trusted-mint.example" - destination_wallet = Mock( - load_proofs=AsyncMock(), - available_balance=Mock(amount=0), - mint=AsyncMock(), - ) - mint_quote = Mock(quote="mint-quote", request="lnbc-test-invoice") - calculate_amount = AsyncMock(return_value=5) - - with ( - patch("routstr.wallet.settings.primary_mint", destination_url), - patch("routstr.wallet.settings.primary_mint_unit", "sat"), - patch("routstr.wallet.settings.cashu_mints", [destination_url]), - patch( - "routstr.wallet._calculate_swap_amount", - calculate_amount, - ), - patch( - "routstr.wallet._request_mint_with_fallback", - AsyncMock(return_value=(destination_wallet, destination_url, mint_quote)), - ), - ): - assert await swap_to_trusted_mint(token, source_wallet) == ( - 5, - "sat", - destination_url, - ) - - calculate_call = calculate_amount.await_args - assert calculate_call is not None - resolved = cast(list[Proof], calculate_call.args[5]) - assert resolved[0].id == FULL_V2_ID - assert source_wallet.get_fees_for_proofs.call_args.args[0] is resolved - assert source_wallet.melt.await_args.kwargs["proofs"] is resolved diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 31dd5767..498b57f5 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -1,7 +1,7 @@ """Tests for stale reserved_balance handling (issue #551). Covers: -- pay_for_request stamping reserved_at on billing and child keys +- pay_for_request stamping reserved_at on charged keys - release_stale_reservations sweeper semantics - reset_all_reserved_balances clearing reserved_at - refund endpoint self-healing stale/legacy reservations @@ -16,14 +16,13 @@ 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, select +from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession 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, ) @@ -70,26 +69,6 @@ async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None: assert key.reserved_at >= before -@pytest.mark.asyncio -async def test_pay_for_request_sets_reserved_at_on_child_key( - session: AsyncSession, -) -> None: - parent = ApiKey(hashed_key="parentkey", balance=10_000) - child = ApiKey(hashed_key="childkey", balance=0, parent_key_hash="parentkey") - session.add(parent) - session.add(child) - await session.commit() - - await pay_for_request(child, 1_000, session) - - await session.refresh(parent) - await session.refresh(child) - assert parent.reserved_balance == 1_000 - assert parent.reserved_at is not None - assert child.reserved_balance == 1_000 - assert child.reserved_at is not None - - @pytest.mark.asyncio async def test_revert_clears_reserved_at_when_fully_released( session: AsyncSession, @@ -153,39 +132,6 @@ async def test_release_stale_reservations_releases_old(session: AsyncSession) -> assert key.reserved_at is None -@pytest.mark.asyncio -async def test_targeted_parent_cleanup_releases_child_owned_reservation( - session: AsyncSession, -) -> None: - parent = ApiKey(hashed_key="stale-parent", balance=5_000) - child = ApiKey( - hashed_key="stale-child", parent_key_hash=parent.hashed_key, balance=0 - ) - session.add_all([parent, child]) - await session.commit() - await pay_for_request(child, 1_000, session) - reservation = ( - await session.exec( - select(ReservationRelease).where( - ReservationRelease.key_hash == child.hashed_key - ) - ) - ).one() - reservation.created_at = int(time.time()) - 1_000 - session.add(reservation) - await session.commit() - - released = await release_stale_reservations( - session, max_age_seconds=300, key_hash=parent.hashed_key - ) - - assert released == 1 - await session.refresh(parent) - await session.refresh(child) - assert parent.reserved_balance == 0 - assert child.reserved_balance == 0 - - @pytest.mark.asyncio async def test_release_stale_reservations_keeps_fresh(session: AsyncSession) -> None: key = ApiKey( @@ -254,10 +200,8 @@ async def test_reset_all_reserved_balances_clears_reserved_at( def _refund_patches(refund_token: str = "cashuArefund"): # type: ignore[no-untyped-def] return ( - patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), + patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), ) @@ -278,8 +222,8 @@ async def test_refund_self_heals_stale_reservation(session: AsyncSession) -> Non reserved_at=int(time.time()) - 10_000, ) - p1, p2, p3, p4 = _refund_patches() - with p1, p2, p3, p4: + p1, p2 = _refund_patches() + with p1, p2: result = await refund_wallet_endpoint( authorization="Bearer sk-stalerefund", x_cashu=None, @@ -307,8 +251,8 @@ async def test_refund_self_heals_legacy_null_reserved_at(session: AsyncSession) reserved_at=None, ) - p1, p2, p3, p4 = _refund_patches() - with p1, p2, p3, p4: + p1, p2 = _refund_patches() + with p1, p2: result = await refund_wallet_endpoint( authorization="Bearer sk-legacyrefund", x_cashu=None, @@ -334,8 +278,8 @@ async def test_refund_rejects_recent_reservation(session: AsyncSession) -> None: reserved_at=int(time.time()), ) - p1, p2, p3, p4 = _refund_patches() - with p1, p2, p3, p4: + p1, p2 = _refund_patches() + with p1, p2: with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( authorization="Bearer sk-activerefund", @@ -356,8 +300,8 @@ async def test_refund_without_reservation_still_works(session: AsyncSession) -> reserved_balance=0, ) - p1, p2, p3, p4 = _refund_patches() - with p1, p2, p3, p4: + p1, p2 = _refund_patches() + with p1, p2: result = await refund_wallet_endpoint( authorization="Bearer sk-plainrefund", x_cashu=None, diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 2ae574ab..e70804c1 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -1,8 +1,12 @@ import asyncio +import json from collections.abc import AsyncGenerator +from typing import cast 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 @@ -18,6 +22,7 @@ from routstr.auth import ( ) from routstr.core.db import ApiKey, ReservationRelease from routstr.payment.cost_calculation import MaxCostData +from routstr.payment.models import Architecture, Model, Pricing from routstr.upstream.base import BaseUpstreamProvider @@ -77,44 +82,43 @@ async def test_release_only_owns_its_concurrent_reservation() -> None: @pytest.mark.asyncio -async def test_release_updates_parent_and_child_atomically() -> None: +async def test_release_clears_reservation_aggregates_atomically() -> None: engine = await _engine() - parent = ApiKey(hashed_key="parent", balance=1_000) - child = ApiKey(hashed_key="child", parent_key_hash="parent", balance=0) + key = ApiKey(hashed_key="key", balance=1_000) async with AsyncSession(engine, expire_on_commit=False) as session: - session.add_all([parent, child]) + session.add(key) await session.commit() - await pay_for_request(child, 500, session) - snapshot = await get_reservation_snapshot(child, session) + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) assert await release_reservation(snapshot, session, 500) is True - await session.refresh(parent) - await session.refresh(child) - assert (parent.reserved_balance, child.reserved_balance) == (0, 0) - assert (parent.reserved_at, child.reserved_at) == (None, None) + await session.refresh(key) + assert (key.reserved_balance, key.reserved_at) == (0, None) await engine.dispose() @pytest.mark.asyncio -async def test_release_rolls_back_partial_parent_child_update() -> None: +async def test_release_repairs_partial_aggregate_corruption() -> None: + """An aggregate that no longer holds the reservation must not leave + the durable row active forever: the release rolls the subtraction back and + terminalizes the reservation without touching aggregates.""" engine = await _engine() - parent = ApiKey(hashed_key="parent", balance=1_000) - child = ApiKey(hashed_key="child", parent_key_hash="parent", balance=0) + key = ApiKey(hashed_key="key", balance=1_000) async with AsyncSession(engine, expire_on_commit=False) as session: - session.add_all([parent, child]) + session.add(key) await session.commit() - await pay_for_request(child, 500, session) - snapshot = await get_reservation_snapshot(child, session) - child.reserved_balance = 100 - session.add(child) + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) + key.reserved_balance = 100 + session.add(key) await session.commit() - assert await release_reservation(snapshot, session, 500) is False - await session.refresh(parent) - await session.refresh(child) + assert await release_reservation(snapshot, session, 500) is True + await session.refresh(key) record = await session.get(ReservationRelease, snapshot.release_id) - assert (parent.reserved_balance, child.reserved_balance) == (500, 100) - assert record is not None and record.status == "active" + # Aggregates untouched — legacy cleanup reconciles them when stale. + assert key.reserved_balance == 100 + assert record is not None and record.status == "released" await engine.dispose() @@ -333,6 +337,220 @@ async def test_responses_streaming_releases_and_raises_on_billing_failure( release.assert_awaited_once_with(snapshot, session, 500) +@pytest.mark.asyncio +@pytest.mark.parametrize("api", ["chat", "responses"]) +@pytest.mark.parametrize("finalization_fails", [False, True]) +async def test_partial_remote_protocol_error_finalizes_and_closes_once( + api: str, + finalization_fails: bool, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + yield b'data: {"model":"test","choices":[{"delta":{"content":"hi"}}]}\n\n' + raise httpx.RemoteProtocolError("incomplete chunked read") + + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + 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 + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.rollback = AsyncMock() + session_context = MagicMock() + session_context.__aenter__ = AsyncMock(return_value=session) + session_context.__aexit__ = AsyncMock(return_value=None) + adjust = ( + AsyncMock(side_effect=SQLAlchemyError("database unavailable")) + if finalization_fails + else AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0}) + ) + snapshot = ReservationSnapshot( + release_id=f"{api}-partial-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + release = AsyncMock(return_value=True) + + with ( + patch("routstr.upstream.base.adjust_payment_for_tokens", adjust), + patch("routstr.upstream.base.release_reservation", release), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + if api == "chat": + response = await provider.handle_streaming_chat_completion( + 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( + response=upstream_response, + 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) + ) + + adjust.assert_awaited_once() + if finalization_fails: + session.rollback.assert_awaited_once() + release.assert_awaited_once_with(snapshot, session, 500) + 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( + api: str, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + yield b'data: {"model":"test","choices":[]}\n\n' + raise httpx.RemoteProtocolError("incomplete chunked read") + + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + 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 + snapshot = ReservationSnapshot( + release_id=f"{api}-database-down-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + unavailable_session = MagicMock() + unavailable_session.__aenter__ = AsyncMock( + side_effect=SQLAlchemyError("database unavailable") + ) + unavailable_session.__aexit__ = AsyncMock(return_value=None) + + with patch( + "routstr.upstream.base.create_session", return_value=unavailable_session + ): + if api == "chat": + response = await provider.handle_streaming_chat_completion( + 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( + response=upstream_response, + 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 + + upstream_response.aclose.assert_awaited_once() + client.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_responses_streaming_duplicate_publishes_zero_settled_cost() -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + yield ( + b'data: {"type":"response.completed","response":{"model":"test",' + b'"usage":{"input_tokens":2,"output_tokens":1}}}\n\n' + ) + yield b"data: [DONE]\n\n" + + upstream_response = MagicMock( + status_code=200, + headers={"content-type": "text/event-stream"}, + ) + upstream_response.aiter_bytes = aiter_bytes + key = MagicMock(spec=ApiKey) + key.hashed_key = "responses-duplicate" + key.balance = 10_000 + 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) + cost_data = { + "input_tokens": 2, + "output_tokens": 1, + "input_msats": 1_000, + "output_msats": 500, + "total_msats": 1_500, + "charged_msats": 0, + "total_usd": 0.0001, + } + + with ( + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + AsyncMock(return_value=cost_data), + ), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + response = await provider.handle_streaming_responses_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + ) + chunks = [ + chunk if isinstance(chunk, str) else bytes(chunk).decode() + async for chunk in response.body_iterator + ] + + completed = next( + json.loads(line[6:]) + for line in "".join(chunks).splitlines() + if line.startswith("data: {") + ) + nested_usage = completed["response"]["usage"] + assert nested_usage["cost_sats"] == 0 + assert nested_usage["cost"]["total_msats"] == 0 + assert nested_usage["cost"]["charged_msats"] == 0 + assert nested_usage["cost"]["computed_msats"] == 1_500 + assert completed["cost"]["total_msats"] == 0 + assert completed["cost"]["computed_msats"] == 1_500 + + @pytest.mark.asyncio @pytest.mark.parametrize("via_litellm", [False, True]) @pytest.mark.parametrize( @@ -446,3 +664,114 @@ async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> assert second.reserved_balance == 0 await engine.dispose() + + +@pytest.mark.asyncio +async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat() -> ( + None +): + """A client abort releases the hold after charging only estimated usage. + + Starlette closes the response generator (``aclose``) on disconnect. The + finalizer still has the request and streamed deltas, so it can estimate + usage without converting the reservation ceiling into the charge. + """ + engine = await _engine() + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key", provider_fee=1.0 + ) + + async with AsyncSession(engine, expire_on_commit=False) as session: + key = ApiKey(hashed_key="disconnect-key", balance=1_000) + session.add(key) + await session.commit() + await pay_for_request(key, 500, session) + snapshot = await get_reservation_snapshot(key, session) + + assert snapshot.release_id in auth_module._reservation_heartbeats + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + # A live stream that never sends a usage chunk or [DONE]; the client + # disconnects after the first delta. + yield b'data: {"choices":[{"delta":{"content":"hi"}}]}\n\n' + yield b'data: {"choices":[{"delta":{"content":" there"}}]}\n\n' + + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + upstream_response.aiter_bytes = aiter_bytes + + model = Model( + id="test-model", + name="test-model", + created=0, + description="", + context_length=8_192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing(prompt=0.01, completion=0.02), + sats_pricing=Pricing(prompt=0.01, completion=0.02), + ) + request_body = json.dumps( + {"model": model.id, "messages": [{"role": "user", "content": "hi"}]} + ).encode() + + background_tasks = BackgroundTasks() + try: + with ( + patch( + "routstr.upstream.base.create_session", + side_effect=lambda: AsyncSession(engine, expire_on_commit=False), + ), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + auth_module.adjust_payment_for_tokens, + ), + patch("routstr.upstream.count_tokens._count_with_litellm", return_value=3), + patch( + "routstr.upstream.count_tokens._count_text_with_litellm", + return_value=2, + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", + return_value=5.0e-5, + ), + ): + response = await provider.handle_streaming_chat_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + background_tasks=background_tasks, + model_obj=model, + reservation_snapshot=snapshot, + request_body=request_body, + ) + 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) + + async with AsyncSession(engine, expire_on_commit=False) as session: + final_key = await session.get(ApiKey, "disconnect-key") + record = await session.get(ReservationRelease, snapshot.release_id) + + assert final_key is not None + # The reservation reached a single terminal outcome; funds are not locked. + assert record is not None and record.status in {"charged", "released"} + assert final_key.reserved_balance == 0 + # 3 input tokens × 10 msats + 2 output tokens × 20 msats = 70 msats. + assert final_key.total_spent == 70 + assert final_key.balance == 930 + # The heartbeat is gone — no forever-renewing task on an abandoned request. + assert snapshot.release_id not in auth_module._reservation_heartbeats + await engine.dispose() diff --git a/tests/unit/test_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index 9435a851..36e76765 100644 --- a/tests/unit/test_tinfoil_integration.py +++ b/tests/unit/test_tinfoil_integration.py @@ -11,18 +11,25 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from routstr import proxy as proxy_module +from routstr.core.db import ApiKey from routstr.upstream.ehbp import ( _PROXY_ONLY_HEADERS, + EHBPForwardingTarget, _compute_ehbp_actual_cost, + _is_ehbp_key_config_response, + _passthrough_key_config_response, _prepare_ehbp_upstream_headers, _resolve_ehbp_target_url, _strip_proxy_headers, + forward_ehbp_x_cashu_request, parse_tinfoil_usage_metrics, ) from routstr.upstream.tinfoil import ( TinfoilModel, TinfoilUpstreamProvider, ) +from routstr.upstream.tinfoil_trailer import TrailerResponse # --------------------------------------------------------------------------- # parse_tinfoil_usage_metrics @@ -102,13 +109,29 @@ class TestParseTinfoilUsageMetrics: assert result["prompt_tokens"] == 69 assert result["completion_tokens"] == 20 assert result["total_tokens"] == 89 + assert result["cache_read_input_tokens"] == 64 + assert result["uncached_prompt_tokens"] == 5 assert result["model"] == "kimi-k2-6" + assert "cost_usd" not in result + + def test_with_cached_and_cost_usd(self) -> None: + result = parse_tinfoil_usage_metrics( + "prompt=69,completion=20,total=89," + "cached_prompt_tokens=64,uncached_prompt_tokens=5," + "model=glm-5-2,cost_usd=0.000123456" + ) + assert result is not None + assert result["prompt_tokens"] == 69 + assert result["completion_tokens"] == 20 + assert result["total_tokens"] == 89 + assert result["cache_read_input_tokens"] == 64 + assert result["uncached_prompt_tokens"] == 5 + assert result["cost_usd"] == 0.000123456 + assert result["model"] == "glm-5-2" def test_old_format_still_works(self) -> None: """Headers without the model field (pre-PR #385) still parse.""" - result = parse_tinfoil_usage_metrics( - "prompt=67,completion=42,total=109" - ) + result = parse_tinfoil_usage_metrics("prompt=67,completion=42,total=109") assert result == { "prompt_tokens": 67, "completion_tokens": 42, @@ -240,12 +263,12 @@ class TestResolveEhbpTargetUrl: class TestComputeEhbpActualCost: @pytest.mark.asyncio - async def test_no_usage_falls_back_to_max_cost(self) -> None: + async def test_no_usage_does_not_charge_authorization_ceiling(self) -> None: model_obj = MagicMock() model_obj.id = "llama3-3-70b" model_obj.forwarded_model_id = "llama3-3-70b" result = await _compute_ehbp_actual_cost(None, model_obj, 100_000) - assert result["total_msats"] == 100_000 + assert result["total_msats"] == 0 assert result["input_tokens"] == 0 assert result["output_tokens"] == 0 @@ -285,7 +308,49 @@ class TestComputeEhbpActualCost: assert result["output_msats"] == 20 @pytest.mark.asyncio - async def test_max_cost_data_falls_back(self) -> None: + async def test_cache_fields_propagated(self) -> None: + model_obj = MagicMock() + model_obj.id = "tinfoil-glm-5-2" + model_obj.forwarded_model_id = "glm-5-2" + with patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc: + from routstr.payment.cost_calculation import CostData + + mock_calc.return_value = CostData( + base_msats=0, + input_msats=5, + output_msats=20, + total_msats=25, + total_usd=0.0003, + input_tokens=5, + output_tokens=20, + cache_read_input_tokens=64, + cache_creation_input_tokens=0, + cache_read_msats=12, + cache_creation_msats=0, + ) + result = await _compute_ehbp_actual_cost( + "prompt=69,completion=20,total=89," + "cached_prompt_tokens=64,uncached_prompt_tokens=5," + "model=glm-5-2", + model_obj, + 100_000, + ) + assert result["total_msats"] == 25 + assert result["input_tokens"] == 5 + assert result["output_tokens"] == 20 + assert result["cache_read_input_tokens"] == 64 + assert result["cache_creation_input_tokens"] == 0 + assert result["cache_read_msats"] == 12 + assert result["cache_creation_msats"] == 0 + assert result["total_usd"] == 0.0003 + + @pytest.mark.asyncio + async def test_unpriceable_usage_does_not_charge_authorization_ceiling( + self, + ) -> None: model_obj = MagicMock() model_obj.id = "llama3-3-70b" model_obj.forwarded_model_id = "llama3-3-70b" @@ -309,7 +374,7 @@ class TestComputeEhbpActualCost: model_obj, 50_000, ) - assert result["total_msats"] == 50_000 + assert result["total_msats"] == 0 assert result["input_tokens"] == 0 assert result["output_tokens"] == 0 @@ -390,13 +455,16 @@ class TestComputeEhbpActualCost: actual_model_obj.id = "tinfoil-llama3-3-70b" # client-facing of actual actual_model_obj.forwarded_model_id = "llama3-3-70b" - with patch( - "routstr.proxy.get_model_instance", - return_value=actual_model_obj, - ), patch( - "routstr.upstream.ehbp.calculate_cost", - new_callable=AsyncMock, - ) as mock_calc: + with ( + patch( + "routstr.proxy.get_model_instance", + return_value=actual_model_obj, + ), + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): from routstr.payment.cost_calculation import CostData mock_calc.return_value = CostData( @@ -420,6 +488,230 @@ class TestComputeEhbpActualCost: call_args = mock_calc.call_args assert call_args[0][0]["model"] == "tinfoil-llama3-3-70b" + @pytest.mark.asyncio + async def test_namespaced_prefix_served_bare_keeps_requested_pricing( + self, + ) -> None: + """Production shape: the catalog model is ``tinfoil-X`` and its + ``forwarded_model_id`` carries the prefix, the SDK strips the prefix + for the encrypted body, and the enclave reports bare ``X``. Pricing + must stay on the requested Tinfoil model (correct rate + cache + discount), not the cheaper cross-provider model the bare id resolves + to in the global map.""" + model_obj = MagicMock() + model_obj.id = "tinfoil-deepseek-v4-1-flash" + model_obj.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + + tinfoil_model = MagicMock() + tinfoil_model.id = "tinfoil-deepseek-v4-1-flash" + tinfoil_model.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + + # The cheaper cross-provider model the bare id resolves to globally. + cross_provider_model = MagicMock() + cross_provider_model.id = "deepseek-v4-1-flash" + cross_provider_model.forwarded_model_id = "deepseek-v4-1-flash" + + registry = { + "tinfoil-deepseek-v4-1-flash": tinfoil_model, + "deepseek-v4-1-flash": cross_provider_model, + } + + with ( + patch( + "routstr.proxy.get_model_instance", + side_effect=lambda name: registry.get(name), + ), + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): + from routstr.payment.cost_calculation import CostData + + mock_calc.return_value = CostData( + base_msats=0, + input_msats=5, + output_msats=10, + total_msats=15, + total_usd=0.0, + input_tokens=5, + output_tokens=10, + cache_read_input_tokens=64, + cache_creation_input_tokens=0, + cache_read_msats=1, + cache_creation_msats=0, + ) + result = await _compute_ehbp_actual_cost( + "prompt=69,completion=10,total=79," + "cached_prompt_tokens=64,uncached_prompt_tokens=5," + "model=deepseek-v4-1-flash", + model_obj, + 100_000, + ) + # No mismatch: pricing stays on the requested Tinfoil model. + assert "actual_model" not in result + call_args = mock_calc.call_args + assert call_args[0][0]["model"] == "tinfoil-deepseek-v4-1-flash" + + @pytest.mark.asyncio + async def test_namespaced_prefix_failover_uses_served_tinfoil_model( + self, + ) -> None: + """A genuine failover (asked ``tinfoil-glm-5-3``, enclave served + ``glm-5-3-flash``) must bill the served *Tinfoil* model, not the + cheaper cross-provider alias the bare id resolves to.""" + model_obj = MagicMock() + model_obj.id = "tinfoil-glm-5-3" + model_obj.forwarded_model_id = "tinfoil-glm-5-3" + + served_tinfoil = MagicMock() + served_tinfoil.id = "tinfoil-glm-5-3-flash" + served_tinfoil.forwarded_model_id = "tinfoil-glm-5-3-flash" + + cross_provider = MagicMock() + cross_provider.id = "glm-5-3-flash" + cross_provider.forwarded_model_id = "glm-5-3-flash" + + registry = { + "tinfoil-glm-5-3-flash": served_tinfoil, + "glm-5-3-flash": cross_provider, + } + + with ( + patch( + "routstr.proxy.get_model_instance", + side_effect=lambda name: registry.get(name), + ), + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): + from routstr.payment.cost_calculation import CostData + + mock_calc.return_value = CostData( + base_msats=0, + input_msats=20, + output_msats=40, + total_msats=60, + total_usd=0.0, + input_tokens=42, + output_tokens=10, + ) + result = await _compute_ehbp_actual_cost( + "prompt=42,completion=10,total=52,model=glm-5-3-flash", + model_obj, + 100_000, + ) + assert result["actual_model"] == "glm-5-3-flash" + # Billed on the served *Tinfoil* model, not the bare-id alias. + call_args = mock_calc.call_args + assert call_args[0][0]["model"] == "tinfoil-glm-5-3-flash" + + @pytest.mark.asyncio + async def test_calculate_cost_receives_routed_model_obj(self) -> None: + """The routed ``Model`` is handed to ``calculate_cost`` so pricing is + billed directly. Without it, ``calculate_cost`` re-derives pricing from + the response's model *string* through the global alias map, which + resolves a bare id to the best-ranked (cheaper) cross-provider + candidate rather than the serving one.""" + model_obj = MagicMock() + model_obj.id = "deepseek-v4-1-flash" + model_obj.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + + resolved = MagicMock() + resolved.id = "tinfoil-deepseek-v4-1-flash" + resolved.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + + with ( + patch( + "routstr.proxy.get_model_instance", + return_value=resolved, + ), + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): + from routstr.payment.cost_calculation import CostData + + mock_calc.return_value = CostData( + base_msats=0, + input_msats=5, + output_msats=10, + total_msats=15, + total_usd=0.0, + input_tokens=5, + output_tokens=10, + ) + await _compute_ehbp_actual_cost( + "prompt=12952,completion=1,total=12953," + "cached_prompt_tokens=12800,uncached_prompt_tokens=152," + "model=deepseek-v4-1-flash,cost_usd=0.00176715", + model_obj, + 100_000, + ) + # The routed model object itself must be passed through. + assert mock_calc.call_args[0][2] is model_obj + + @pytest.mark.asyncio + async def test_routed_model_cache_rate_beats_bare_id_alias(self) -> None: + """Production regression: the routed Tinfoil model's id *is* a bare + cross-provider alias, so re-deriving pricing from the echoed model + string silently swapped in the cheaper candidate's full input rate and + the cache discount vanished. Billing must use the routed model's own + discounted cache rate.""" + from routstr.payment.models import Pricing + + model_obj = MagicMock() + model_obj.id = "deepseek-v4-1-flash" + model_obj.forwarded_model_id = "tinfoil-deepseek-v4-1-flash" + # ~688 msat/1k input, ~112 msat/1k cached read (the good Tinfoil rate). + model_obj.sats_pricing = Pricing( + prompt=6.88e-4, + completion=2.0e-3, + input_cache_read=1.12e-4, + ) + + # The cross-provider candidate the bare id resolves to globally: no + # cache rate at all, so a re-derivation charges the full input rate. + cross_provider_model = MagicMock() + cross_provider_model.id = "deepseek-v4-1-flash" + cross_provider_model.forwarded_model_id = "deepseek-v4-1-flash" + cross_provider_model.sats_pricing = Pricing( + prompt=4.9455e-4, + completion=2.0e-3, + input_cache_read=0.0, + ) + + registry = { + "tinfoil-deepseek-v4-1-flash": model_obj, + "deepseek-v4-1-flash": cross_provider_model, + } + + with ( + patch( + "routstr.proxy.get_model_instance", + side_effect=lambda name: registry.get(name), + ), + patch( + "routstr.payment.cost_calculation.sats_usd_price", + return_value=5.0e-5, + ), + ): + result = await _compute_ehbp_actual_cost( + "prompt=12952,completion=1,total=12953," + "cached_prompt_tokens=12800,uncached_prompt_tokens=152," + "model=deepseek-v4-1-flash,cost_usd=0.00176715", + model_obj, + 100_000, + ) + + assert result["cache_read_input_tokens"] == 12800 + # 12800 cached tokens at the discounted (~112 msat/1k) rate, not the + # full input rate (which would be ~8800 msat here). + assert result["cache_read_msats"] == pytest.approx(1434, abs=10) + @pytest.mark.asyncio async def test_model_mismatch_unknown_model_falls_back(self) -> None: """When the served model is not in the registry, use requested model.""" @@ -427,13 +719,16 @@ class TestComputeEhbpActualCost: model_obj.id = "gpt-oss-120b" model_obj.forwarded_model_id = "gpt-oss-120b" - with patch( - "routstr.proxy.get_model_instance", - return_value=None, - ), patch( - "routstr.upstream.ehbp.calculate_cost", - new_callable=AsyncMock, - ) as mock_calc: + with ( + patch( + "routstr.proxy.get_model_instance", + return_value=None, + ), + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): from routstr.payment.cost_calculation import CostData mock_calc.return_value = CostData( @@ -492,12 +787,13 @@ class TestComputeEhbpActualCost: model_obj = MagicMock() model_obj.id = "tinfoil-glm-5-2" model_obj.forwarded_model_id = "glm-5-2" # lowercase - with patch( - "routstr.proxy.get_model_instance" - ) as mock_get_model, patch( - "routstr.upstream.ehbp.calculate_cost", - new_callable=AsyncMock, - ) as mock_calc: + with ( + patch("routstr.proxy.get_model_instance") as mock_get_model, + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): from routstr.payment.cost_calculation import CostData mock_calc.return_value = CostData( @@ -533,13 +829,16 @@ class TestComputeEhbpActualCost: resolved_model_obj.id = "other-provider-glm-5-2" resolved_model_obj.forwarded_model_id = "glm-5-2" - with patch( - "routstr.proxy.get_model_instance", - return_value=resolved_model_obj, - ) as mock_get_model, patch( - "routstr.upstream.ehbp.calculate_cost", - new_callable=AsyncMock, - ) as mock_calc: + with ( + patch( + "routstr.proxy.get_model_instance", + return_value=resolved_model_obj, + ) as mock_get_model, + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): from routstr.payment.cost_calculation import CostData mock_calc.return_value = CostData( @@ -570,12 +869,13 @@ class TestComputeEhbpActualCost: model_obj.id = "tinfoil-glm-5-2-20260415" model_obj.forwarded_model_id = "glm-5-2-20260415" - with patch( - "routstr.proxy.get_model_instance" - ) as mock_get_model, patch( - "routstr.upstream.ehbp.calculate_cost", - new_callable=AsyncMock, - ) as mock_calc: + with ( + patch("routstr.proxy.get_model_instance") as mock_get_model, + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): from routstr.payment.cost_calculation import CostData mock_calc.return_value = CostData( @@ -594,10 +894,7 @@ class TestComputeEhbpActualCost: ) assert "actual_model" not in result - assert ( - mock_calc.call_args[0][0]["model"] - == "tinfoil-glm-5-2-20260415" - ) + assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2-20260415" mock_get_model.assert_not_called() @pytest.mark.asyncio @@ -612,13 +909,16 @@ class TestComputeEhbpActualCost: resolved_model_obj.id = "other-provider-glm-5-2" resolved_model_obj.forwarded_model_id = "GLM-5-2" - with patch( - "routstr.proxy.get_model_instance", - return_value=resolved_model_obj, - ) as mock_get_model, patch( - "routstr.upstream.ehbp.calculate_cost", - new_callable=AsyncMock, - ) as mock_calc: + with ( + patch( + "routstr.proxy.get_model_instance", + return_value=resolved_model_obj, + ) as mock_get_model, + patch( + "routstr.upstream.ehbp.calculate_cost", + new_callable=AsyncMock, + ) as mock_calc, + ): from routstr.payment.cost_calculation import CostData mock_calc.return_value = CostData( @@ -650,8 +950,7 @@ class TestTinfoilUpstreamProvider: def test_provider_type_and_defaults(self) -> None: assert TinfoilUpstreamProvider.provider_type == "tinfoil" assert ( - TinfoilUpstreamProvider.default_base_url - == "https://inference.tinfoil.sh" + TinfoilUpstreamProvider.default_base_url == "https://inference.tinfoil.sh" ) assert TinfoilUpstreamProvider.supports_ehbp is True @@ -666,9 +965,7 @@ class TestTinfoilUpstreamProvider: model_obj.id = "llama3-3-70b" model_obj.forwarded_model_id = "llama3-3-70b" target = provider.get_ehbp_forwarding_target("v1/chat/completions", model_obj) - assert ( - target.headers["X-Tinfoil-Request-Usage-Metrics"] == "true" - ) + assert target.headers["X-Tinfoil-Request-Usage-Metrics"] == "true" assert "v1/chat/completions" in target.url def test_get_provider_metadata(self) -> None: @@ -694,6 +991,20 @@ class TestTinfoilUpstreamProvider: assert tf.id == "llama3-3-70b" assert tf.pricing.inputTokenPricePer1M == 1.75 assert tf.pricing.outputTokenPricePer1M == 2.75 + assert tf.pricing.cachedInputTokenPricePer1M is None + + def test_tinfoil_model_pricing_parses_cached_rate(self) -> None: + data = { + "id": "glm-5-2", + "pricing": { + "inputTokenPricePer1M": 1.5, + "outputTokenPricePer1M": 5.25, + "cachedInputTokenPricePer1M": 0.375, + "requestPrice": 0, + }, + } + tf = TinfoilModel.parse_obj(data) + assert tf.pricing.cachedInputTokenPricePer1M == 0.375 @pytest.mark.asyncio async def test_fetch_models_parses_response(self) -> None: @@ -732,8 +1043,52 @@ class TestTinfoilUpstreamProvider: assert models[0].id == "llama3-3-70b" assert models[0].pricing.prompt == 1.75 / 1_000_000 assert models[0].pricing.completion == 2.75 / 1_000_000 + # No cachedInputTokenPricePer1M means cache reads are billed at the + # full input rate (and cache writes too — no separate write price). + assert models[0].pricing.input_cache_read == 1.75 / 1_000_000 + assert models[0].pricing.input_cache_write == 1.75 / 1_000_000 assert models[0].context_length == 128000 + @pytest.mark.asyncio + async def test_fetch_models_maps_cached_pricing(self) -> None: + provider = TinfoilUpstreamProvider(api_key="test") + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.raise_for_status = MagicMock() + mock_response.json.return_value = { + "data": [ + { + "id": "glm-5-2", + "context_window": 393216, + "created": 1775088000, + "multimodal": False, + "pricing": { + "inputTokenPricePer1M": 1.5, + "outputTokenPricePer1M": 5.25, + "cachedInputTokenPricePer1M": 0.375, + "requestPrice": 0, + }, + "endpoints": ["/v1/chat/completions", "/v1/responses"], + "type": "chat", + } + ] + } + + with patch("routstr.upstream.tinfoil.httpx.AsyncClient") as mock_client_cls: + mock_client = MagicMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.get = AsyncMock(return_value=mock_response) + mock_client_cls.return_value = mock_client + + models = await provider.fetch_models() + + assert len(models) == 1 + assert models[0].pricing.prompt == 1.5 / 1_000_000 + assert models[0].pricing.completion == 5.25 / 1_000_000 + assert models[0].pricing.input_cache_read == 0.375 / 1_000_000 + assert models[0].pricing.input_cache_write == 1.5 / 1_000_000 + @pytest.mark.asyncio async def test_fetch_models_handles_error(self) -> None: provider = TinfoilUpstreamProvider(api_key="test") @@ -747,3 +1102,276 @@ class TestTinfoilUpstreamProvider: models = await provider.fetch_models() assert models == [] + + +# --------------------------------------------------------------------------- +# EHBP key-config mismatch passthrough +# --------------------------------------------------------------------------- + + +def _key_config_trailer_response( + status_code: int = 422, + content_type: str = "application/problem+json", + body: bytes = b'{"type":"urn:ietf:params:ehbp:error:key-config","title":"failed to read decrypted request body"}', +) -> TrailerResponse: + return TrailerResponse( + status_code=status_code, + headers=[ + ("content-type", content_type), + ("content-length", str(len(body))), + ], + body=body, + trailers=[], + ) + + +class TestIsEhbpKeyConfigResponse: + def test_genuine_key_config_422(self) -> None: + resp = _key_config_trailer_response() + assert _is_ehbp_key_config_response(resp) is True + + def test_200_is_not_key_config(self) -> None: + resp = _key_config_trailer_response(status_code=200) + assert _is_ehbp_key_config_response(resp) is False + + def test_400_is_not_key_config(self) -> None: + resp = _key_config_trailer_response(status_code=400) + assert _is_ehbp_key_config_response(resp) is False + + def test_json_content_type_is_not_key_config(self) -> None: + """A 422 with application/json (e.g. proxy-wrapped error) must NOT be + treated as key-config — only the original problem+json counts.""" + resp = _key_config_trailer_response(content_type="application/json") + assert _is_ehbp_key_config_response(resp) is False + + def test_problem_json_with_different_type_is_not_key_config(self) -> None: + """A 422 problem+json with a different error type is not key-config.""" + body = b'{"type":"urn:ietf:params:ehbp:error:other","title":"other"}' + resp = _key_config_trailer_response(body=body) + assert _is_ehbp_key_config_response(resp) is False + + def test_empty_body_is_not_key_config(self) -> None: + resp = _key_config_trailer_response(body=b"") + assert _is_ehbp_key_config_response(resp) is False + + def test_invalid_json_body_is_not_key_config(self) -> None: + resp = _key_config_trailer_response(body=b"not json") + assert _is_ehbp_key_config_response(resp) is False + + def test_problem_json_with_charset(self) -> None: + resp = _key_config_trailer_response( + content_type="application/problem+json; charset=utf-8" + ) + assert _is_ehbp_key_config_response(resp) is True + + def test_content_type_parameter_disguising_other_media_type(self) -> None: + """A substring check would accept this; the media type must match + exactly, mirroring the ehbp client's isProblemJSONContentType.""" + resp = _key_config_trailer_response( + content_type="text/html; x=application/problem+json" + ) + assert _is_ehbp_key_config_response(resp) is False + + def test_uppercase_media_type_with_params_matches(self) -> None: + resp = _key_config_trailer_response( + content_type="Application/Problem+JSON; charset=UTF-8" + ) + assert _is_ehbp_key_config_response(resp) is True + + def test_missing_content_type_is_not_key_config(self) -> None: + resp = TrailerResponse( + status_code=422, + headers=[], + body=b'{"type":"urn:ietf:params:ehbp:error:key-config"}', + ) + assert _is_ehbp_key_config_response(resp) is False + + +class TestPassthroughKeyConfigResponse: + def test_status_and_content_type(self) -> None: + resp = _key_config_trailer_response() + result = _passthrough_key_config_response(resp) + assert result.status_code == 422 + assert result.media_type == "application/problem+json" + + def test_body_passed_through(self) -> None: + original_body = b'{"type":"urn:ietf:params:ehbp:error:key-config","title":"failed to read decrypted request body"}' + resp = _key_config_trailer_response(body=original_body) + result = _passthrough_key_config_response(resp) + assert result.body == original_body + + def test_ehbp_nonce_header_dropped(self) -> None: + """The nonce must not survive the passthrough: the stock ehbp client + checks for the nonce before the key-config mismatch, so a forwarded + nonce would send it down the decrypt path on this plaintext error + body and the re-attestation loop would never fire.""" + resp = TrailerResponse( + status_code=422, + headers=[ + ("content-type", "application/problem+json"), + ("ehbp-response-nonce", "abc123"), + ("content-length", "999"), + ("server", "nginx"), + ("x-request-id", "some-id"), + ], + body=b'{"type":"urn:ietf:params:ehbp:error:key-config","title":"test"}', + ) + result = _passthrough_key_config_response(resp) + assert "ehbp-response-nonce" not in result.headers + assert "server" not in result.headers + assert "x-request-id" not in result.headers + # Content-length is recomputed from the actual body, not forwarded. + assert result.headers["content-length"] == str(len(resp.body)) + + +# --------------------------------------------------------------------------- +# Key-config passthrough at the forwarding call sites +# --------------------------------------------------------------------------- + + +def _ehbp_tinfoil_upstream() -> MagicMock: + """A minimal EHBP-capable upstream stub shaped like the Tinfoil provider.""" + upstream = MagicMock() + upstream.provider_type = "tinfoil" + upstream.supports_ehbp = True + upstream.prepare_headers = MagicMock(side_effect=lambda h: h) + upstream.get_confidential_inference_profile = MagicMock(return_value=None) + upstream.get_ehbp_forwarding_target = MagicMock( + return_value=EHBPForwardingTarget( + url="https://inference.tinfoil.sh/private/v1/chat/completions" + ) + ) + upstream.prepare_params = MagicMock(return_value={}) + return upstream + + +@pytest.mark.asyncio +async def test_bearer_key_config_422_releases_reservation_and_passes_through() -> None: + """The bearer path returns the enclave's problem+json verbatim AND the + reservation is released. + + The early return inside ``forward_ehbp_request`` skips the UpstreamError + handler, so the release depends on the proxy's non-200 branch treating 422 + as non-retryable. Nothing else pins that; this does. + """ + key = ApiKey(hashed_key="keyconfig", balance=10_000) + session = MagicMock() + reservation_snapshot = MagicMock() + revert_mock = AsyncMock(return_value=True) + + request = MagicMock() + request.method = "POST" + request.headers = { + "authorization": "Bearer sk-keyconfig", + "ehbp-encapsulated-key": "abc123", + "x-routstr-model": "tinfoil/llama3-3-70b", + } + request.body = AsyncMock(return_value=b"sealed-body") + request.query_params = {} + + model_obj = MagicMock() + model_obj.id = "tinfoil/llama3-3-70b" + upstream = _ehbp_tinfoil_upstream() + + # The enclave may include a nonce even on the 422 — the passthrough must + # drop it, or stock ehbp clients (nonce checked before key-config) would + # try to decrypt this plaintext body instead of re-attesting. + upstream_resp = _key_config_trailer_response() + upstream_resp.headers.append(("ehbp-response-nonce", "nonce-value")) + + with ( + patch.object( + proxy_module, "get_candidates", return_value=[(model_obj, upstream)] + ), + 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=1_000)), + patch.object( + proxy_module, + "get_reservation_snapshot", + AsyncMock(return_value=reservation_snapshot), + ), + patch.object(proxy_module, "revert_pay_for_request", revert_mock), + patch( + "routstr.upstream.ehbp.forward_with_trailer", + AsyncMock(return_value=upstream_resp), + ), + ): + response = await proxy_module.proxy( + request, "v1/chat/completions", session=session + ) + + # The reservation was released despite the early passthrough return. + revert_mock.assert_awaited_once_with(key, session, 1_000, reservation_snapshot) + # The client receives the enclave's problem+json verbatim... + assert response.status_code == 422 + assert response.headers["content-type"] == "application/problem+json" + assert response.body == upstream_resp.body + # ...without the nonce. + assert "ehbp-response-nonce" not in response.headers + + +@pytest.mark.asyncio +async def test_x_cashu_key_config_422_refunds_and_sets_x_cashu_header() -> None: + """The x-cashu path refunds the full redeemed amount and attaches the + refund token to the passthrough response.""" + request = MagicMock() + request.method = "POST" + request.headers = { + "ehbp-encapsulated-key": "abc123", + "x-routstr-model": "tinfoil/llama3-3-70b", + } + request.query_params = {} + request.body = AsyncMock(return_value=b"sealed-body") + request.state.request_id = "req-1" + + model_obj = MagicMock() + model_obj.id = "tinfoil/llama3-3-70b" + upstream = _ehbp_tinfoil_upstream() + + upstream_resp = _key_config_trailer_response() + refund_mock = AsyncMock(return_value="cashuArefund") + store_mock = AsyncMock() + + with ( + patch( + "routstr.upstream.ehbp.recieve_token", + AsyncMock(return_value=(50_000, "msat", "https://mint.example")), + ), + patch("routstr.upstream.ehbp.store_cashu_transaction", store_mock), + patch("routstr.upstream.ehbp.send_cashu_refund", refund_mock), + patch( + "routstr.upstream.ehbp.forward_with_trailer", + AsyncMock(return_value=upstream_resp), + ), + ): + response = await forward_ehbp_x_cashu_request( + request=request, + x_cashu_token="cashuAtoken", + path="v1/chat/completions", + max_cost_for_model=1_000, + model_obj=model_obj, + upstream=upstream, + ) + + # Full refund of the redeemed amount (the enclave never processed it). + refund_mock.assert_awaited_once_with( + 50_000, "msat", "https://mint.example", "req-1" + ) + # The redemption itself was recorded. + store_mock.assert_awaited_once() + assert store_mock.await_args is not None + assert store_mock.await_args.kwargs.get("typ") == "in" + # Passthrough shape with the refund attached. + assert response.status_code == 422 + assert response.headers["content-type"] == "application/problem+json" + assert response.body == upstream_resp.body + assert response.headers["x-cashu"] == "cashuArefund" diff --git a/tests/unit/test_tinfoil_trailer.py b/tests/unit/test_tinfoil_trailer.py index ef1c96f1..3e4d3e0f 100644 --- a/tests/unit/test_tinfoil_trailer.py +++ b/tests/unit/test_tinfoil_trailer.py @@ -1,9 +1,11 @@ from __future__ import annotations +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest +from routstr.core.exceptions import EhbpTimeoutError, UpstreamError from routstr.upstream.tinfoil_trailer import forward_with_trailer @@ -28,6 +30,14 @@ class FakeWriter: self.written += data +class HangingReader: + """A reader that never returns data, used to trigger a read timeout.""" + + async def read(self, _size: int) -> bytes: + await asyncio.sleep(3600) + return b"" + + @pytest.mark.asyncio async def test_forward_with_trailer_captures_usage_trailer( monkeypatch: pytest.MonkeyPatch, @@ -136,3 +146,62 @@ async def test_forward_with_trailer_enforces_response_size_limit( ) writer.close.assert_called_once() + + +@pytest.mark.asyncio +async def test_forward_with_trailer_connect_timeout_raises_ehbp_timeout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def _hang_connect(*_args: object, **_kwargs: object) -> object: + raise asyncio.TimeoutError + + monkeypatch.setattr( + "routstr.upstream.tinfoil_trailer.asyncio.open_connection", _hang_connect + ) + + with pytest.raises(EhbpTimeoutError, match="connecting"): + await forward_with_trailer( + method="POST", + url="https://enclave.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, +) -> None: + reader = HangingReader() + writer = FakeWriter() + monkeypatch.setattr( + "routstr.upstream.tinfoil_trailer.asyncio.open_connection", + AsyncMock(return_value=(reader, writer)), + ) + + with pytest.raises(EhbpTimeoutError, match="waiting for response data"): + await forward_with_trailer( + method="POST", + url="https://enclave.tinfoil.sh/v1/chat/completions", + headers={}, + body=b"opaque", + timeout_seconds=0.01, + ) + + writer.close.assert_called_once() + + +def test_ehbp_timeout_error_metadata() -> None: + exc = EhbpTimeoutError("boom") + assert exc.status_code == 504 + assert exc.code == "UPSTREAM_TIMEOUT" + assert exc.details is None + assert isinstance(exc, UpstreamError) + + +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.code == "UPSTREAM_TIMEOUT" diff --git a/tests/unit/test_upstream_generic.py b/tests/unit/test_upstream_generic.py index 39daff5c..2e6fd56e 100644 --- a/tests/unit/test_upstream_generic.py +++ b/tests/unit/test_upstream_generic.py @@ -522,3 +522,205 @@ async def test_unresolvable_model_fails_closed( for rec in caplog.records if rec.levelno >= logging.WARNING ) + + +# --------------------------------------------------------------------------- +# rate validation — a malformed rate is not a resolved price, at any rung +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_non_finite_litellm_rate_is_not_a_resolved_price() -> None: + """A non-finite entry in the cost map must not answer the resolution chain. + + The litellm rung rejects negatives and a both-zero entry, but every + comparison with ``NaN`` is False and ``inf`` reads as a large positive, so + either would be reported as a resolved price — enabling the model at a rate + the node cannot bill on. Fail closed instead: the model imports disabled, + which is what "no source knows this price" already means here. + """ + payload = { + "data": [ + {"id": "nan-priced-model", "object": "model", "owned_by": "mystery"}, + ] + } + cost_entry = { + "input_cost_per_token": float("nan"), + "output_cost_per_token": float("inf"), + "max_input_tokens": 8192, + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch("routstr.payment.models.litellm_cost_entry", lambda _id: cost_entry): + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider( + base_url="http://x" + ).fetch_models() + + model = _model_by_id(models, "nan-priced-model") + assert model.enabled is False + assert model.pricing.prompt == 0.0 + assert model.pricing.completion == 0.0 + + +@pytest.mark.asyncio +async def test_non_finite_openrouter_rate_is_not_a_resolved_price() -> None: + """``float("Infinity")`` parses happily from a feed string, so the + OpenRouter rung's coercion accepted it and reported an infinite rate as a + resolved price. It is not a price; the model must import disabled.""" + payload = { + "data": [ + {"id": "or-nonfinite-xyz", "object": "model", "owned_by": "mystery"}, + ] + } + feed = [ + { + "id": "or-nonfinite-xyz", + "pricing": {"prompt": "Infinity", "completion": "0.000002"}, + "context_length": 8192, + } + ] + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=feed) + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "or-nonfinite-xyz") + assert model.enabled is False + assert model.pricing.prompt == 0.0 + assert model.pricing.completion == 0.0 + + +@pytest.mark.asyncio +async def test_non_finite_openrouter_cache_rate_is_dropped_not_carried() -> None: + """A malformed *cache* rate must cost the cache rate, not the model. + + The catalog import filter only inspects prompt and completion, so an entry + with two sound token rates and an unusable ``input_cache_read`` reaches the + resolver intact. Carrying that rate through would price every cached input + token at ``inf``; dropping it falls back to the full input rate, which is + what a missing cache rate already means. + """ + payload = { + "data": [ + {"id": "or-badcache-xyz", "object": "model", "owned_by": "mystery"}, + ] + } + feed = [ + { + "id": "or-badcache-xyz", + "pricing": { + "prompt": "0.000001", + "completion": "0.000002", + "input_cache_read": "Infinity", + }, + "context_length": 8192, + } + ] + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=feed) + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "or-badcache-xyz") + assert model.enabled is True + assert model.pricing.prompt == pytest.approx(1e-06) + assert model.pricing.input_cache_read == 0.0 + + +@pytest.mark.asyncio +async def test_negative_openrouter_cache_rate_is_dropped_not_carried() -> None: + """A negative rate is as unusable as a non-finite one, and arrives the same way. + + The coercion that reads a feed price rejected ``NaN``/``inf`` but returned a + negative unchanged, and the catalog import filter that would have caught one + inspects only prompt and completion. So a negative cache rate was the single + malformed value that still reached a stored price — where it prices cached + input tokens at a credit rather than a charge. + """ + payload = { + "data": [ + {"id": "or-negcache-xyz", "object": "model", "owned_by": "mystery"}, + ] + } + feed = [ + { + "id": "or-negcache-xyz", + "pricing": { + "prompt": "0.000001", + "completion": "0.000002", + "input_cache_read": "-0.0000005", + }, + "context_length": 8192, + } + ] + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=feed) + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "or-negcache-xyz") + assert model.enabled is True + assert model.pricing.prompt == pytest.approx(1e-06) + assert model.pricing.input_cache_read == 0.0 + + +@pytest.mark.asyncio +async def test_boolean_litellm_rate_is_not_a_resolved_price() -> None: + """``isinstance(True, int)`` is True, so a boolean passed the cost map's own + numeric check and resolved as a rate of ``1.0`` — a dollar per token.""" + payload = { + "data": [ + {"id": "bool-priced-model", "object": "model", "owned_by": "mystery"}, + ] + } + cost_entry = { + "input_cost_per_token": True, + "output_cost_per_token": 2e-06, + "max_input_tokens": 8192, + } + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=[]) + with patch("routstr.payment.models.litellm_cost_entry", lambda _id: cost_entry): + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider( + base_url="http://x" + ).fetch_models() + + model = _model_by_id(models, "bool-priced-model") + assert model.enabled is False + assert model.pricing.prompt == 0.0 + assert model.pricing.completion == 0.0 + + +@pytest.mark.asyncio +async def test_boolean_openrouter_rate_is_not_a_resolved_price() -> None: + """The same coercion reads a feed's ``true`` as a rate of ``1.0``; the model + must import disabled rather than priced at a dollar per token.""" + payload = { + "data": [ + {"id": "or-bool-xyz", "object": "model", "owned_by": "mystery"}, + ] + } + feed = [ + { + "id": "or-bool-xyz", + "pricing": {"prompt": True, "completion": "0.000002"}, + "context_length": 8192, + } + ] + + with _patch_models_endpoint(payload): + or_feed = AsyncMock(return_value=feed) + with patch("routstr.payment.models.async_fetch_openrouter_models", or_feed): + models = await GenericUpstreamProvider(base_url="http://x").fetch_models() + + model = _model_by_id(models, "or-bool-xyz") + assert model.enabled is False + assert model.pricing.prompt == 0.0 + assert model.pricing.completion == 0.0 diff --git a/tests/unit/test_upstream_reported_cost.py b/tests/unit/test_upstream_reported_cost.py new file mode 100644 index 00000000..0ee7f9d0 --- /dev/null +++ b/tests/unit/test_upstream_reported_cost.py @@ -0,0 +1,123 @@ +"""Tests that the upstream's own USD figure survives the provider-fee multiply. + +``_calculate_from_usd_cost`` multiplies the fee into the same local that holds +the upstream's reported cost, so by the time a ``CostData`` exists the raw +figure is gone and only the marked-up one remains. Dividing the total back out +does not recover it: when the request falls through to token pricing the total +is the node's own arithmetic, and dividing it compares that number to itself. + +These tests cover ``upstream_usd`` — the pre-fee figure carried alongside the +billed one, and the discriminator for whether an upstream reported a cost at +all. They also cover the boundary it must not cross: the pair of numbers spells +out the node's margin, so it stays internal and is never serialised to a client. +""" + +from __future__ import annotations + +import math +from collections.abc import Iterator +from typing import Any +from unittest.mock import patch + +import pytest + +from routstr.payment.cost_calculation import CostData, calculate_cost +from routstr.payment.models import Architecture, Model, Pricing + + +@pytest.fixture(autouse=True) +def patch_sats_usd_price() -> Iterator[None]: + """Pin the exchange rate; these tests are about the USD figure, not the feed.""" + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5): + yield + + +def _model() -> Model: + return Model( + id="m", + name="m", + created=0, + description="d", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing(prompt=1e-06, completion=2e-06), + sats_pricing=Pricing(prompt=1e-06, completion=2e-06), + ) + + +def _response(usage: dict[str, Any]) -> dict[str, Any]: + return {"model": "m", "usage": usage} + + +@pytest.mark.asyncio +async def test_reported_cost_is_kept_alongside_the_billed_one() -> None: + """The billed total carries the fee; ``upstream_usd`` must not.""" + response = _response( + {"prompt_tokens": 1000, "completion_tokens": 500, "cost": 0.01} + ) + + cost = await calculate_cost( + response, max_cost=999999, model_obj=_model(), provider_fee=1.05 + ) + + assert isinstance(cost, CostData) + assert cost.total_usd == pytest.approx(0.0105) + assert cost.upstream_usd == pytest.approx(0.01) + + +@pytest.mark.asyncio +async def test_token_priced_request_reports_no_upstream_cost() -> None: + """Nothing was reported, so there is nothing to carry — not our own total.""" + response = _response({"prompt_tokens": 1000, "completion_tokens": 500}) + + cost = await calculate_cost( + response, max_cost=999999, model_obj=_model(), provider_fee=1.05 + ) + + assert isinstance(cost, CostData) + assert cost.total_msats > 0 + assert cost.upstream_usd == 0.0 + + +@pytest.mark.asyncio +async def test_reported_cost_is_never_serialised_to_a_client() -> None: + """Publishing it beside the billed total would spell out the node's margin.""" + response = _response( + {"prompt_tokens": 1000, "completion_tokens": 500, "cost": 0.01} + ) + + cost = await calculate_cost( + response, max_cost=999999, model_obj=_model(), provider_fee=1.05 + ) + + assert isinstance(cost, CostData) + assert cost.upstream_usd == pytest.approx(0.01) + assert "upstream_usd" not in cost.dict() + assert "upstream_usd" not in cost.json() + + +@pytest.mark.asyncio +async def test_billed_total_is_reproducible_from_the_reported_cost() -> None: + """Fee, rate and rounding must carry the reported figure to the billed one. + + The identity a report can restate: whatever the upstream said, times the + provider fee, converted at the current rate and rounded up, is what the node + charged. A figure chosen for its awkward remainder keeps the ceiling honest. + """ + response = _response( + {"prompt_tokens": 1000, "completion_tokens": 500, "cost": 0.000123} + ) + + cost = await calculate_cost( + response, max_cost=999999, model_obj=_model(), provider_fee=1.03 + ) + + assert isinstance(cost, CostData) + assert cost.total_msats == math.ceil(cost.upstream_usd * 1.03 / 5.0e-5 * 1000) + assert cost.total_msats == 2534 diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 7c5de1fe..86c90b63 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -8,7 +8,6 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx import pytest -from cashu.core.base import MeltQuoteState from routstr.core.db import ApiKey from routstr.wallet import ( @@ -16,6 +15,7 @@ from routstr.wallet import ( Bolt11PaymentNotAttempted, MintConnectionError, TokenConsumedError, + UntrustedSourceMintError, _is_mint_rate_limited, classify_redemption_error, credit_balance, @@ -26,6 +26,7 @@ from routstr.wallet import ( recieve_token, send, send_token, + send_token_from_owner_locked, ) @@ -40,13 +41,19 @@ def isolate_wallet_runtime_state() -> Generator[None, None, None]: 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_load_locks.clear() + wallet_module._mint_metadata_last_load.clear() + wallet_module._mint_metadata_load_locks.clear() yield settings.mint_max_concurrency = original_concurrency 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_load_locks.clear() + wallet_module._mint_metadata_last_load.clear() + wallet_module._mint_metadata_load_locks.clear() @pytest.mark.asyncio @@ -66,6 +73,63 @@ async def test_get_balance() -> None: assert balance == 50000 +@pytest.mark.asyncio +async def test_wallet_metadata_is_reused_across_units() -> None: + from routstr.wallet import Wallet + + sat_wallet = MagicMock(url="http://mint:3338") + sat_wallet.load_mint_keysets = AsyncMock() + sat_wallet.activate_keyset = AsyncMock() + sat_wallet.load_mint_info = AsyncMock() + sat_wallet.load_keysets_from_db = AsyncMock() + + msat_wallet = MagicMock(url="http://mint:3338") + msat_wallet.load_mint_keysets = AsyncMock() + msat_wallet.activate_keyset = AsyncMock() + msat_wallet.load_mint_info = AsyncMock() + msat_wallet.load_keysets_from_db = AsyncMock() + + with patch("routstr.wallet.time.monotonic", return_value=1000.0): + await Wallet.load_mint(sat_wallet) + await Wallet.load_mint(msat_wallet) + + sat_wallet.load_mint_keysets.assert_awaited_once_with(False) + sat_wallet.load_mint_info.assert_awaited_once_with(reload=True) + msat_wallet.load_mint_keysets.assert_not_awaited() + msat_wallet.load_keysets_from_db.assert_awaited_once_with() + msat_wallet.load_mint_info.assert_awaited_once_with(reload=False) + + +@pytest.mark.asyncio +async def test_get_wallet_refreshes_local_proofs_without_reloading_mint() -> None: + from routstr import wallet as wallet_module + 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)), + patch("routstr.wallet.time.monotonic", return_value=1000.0), + ): + await get_wallet("http://mint:3338", "sat") + wallet_module._wallet_last_load["http://mint:3338_sat"] = 900.0 + await get_wallet("http://mint:3338", "sat") + + assert mock_wallet.load_mint.await_count == 1 + assert mock_wallet.load_proofs.await_count == 2 + + +@pytest.mark.asyncio +async def test_get_wallet_quote_only_skips_proof_reload() -> 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", load_proofs=False) + + mock_wallet.load_mint.assert_awaited_once_with() + mock_wallet.load_proofs.assert_not_awaited() + + @pytest.mark.asyncio async def test_get_wallet_force_reload_bypasses_reload_interval() -> None: from routstr.wallet import get_wallet @@ -159,7 +223,7 @@ async def test_recieve_token_trusted_mint_deducts_input_fee() -> None: """A trusted mint that charges NUT-02 input fees. The same-mint receive (`wallet.split(..., include_fees=True)`, a NUT-03 swap - at the same mint — not swap_to_primary_mint) pays the mint's per-proof fee, + at the same mint) pays the mint's per-proof fee, so routstr only ends up with `face - input_fee` in fresh proofs. The credited amount must reflect that, otherwise routstr over-credits the user and its own wallet drifts toward insolvency. @@ -220,11 +284,13 @@ async def test_recieve_token_trusted_mint_deducts_input_fee() -> None: @pytest.mark.asyncio -async def test_recieve_token_uses_only_requested_destination_mint() -> None: +async def test_recieve_token_redeems_on_issuing_mint_never_swaps() -> None: + """A token from a secondary trusted mint stays on that mint even though a + different primary mint is configured; no cross-mint swap is attempted.""" from routstr.core.settings import settings - source = "http://foreign:3338" - destination = "http://key-mint:3338" + source = "http://secondary:3338" + primary = "http://primary:3338" token = Mock( mint=source, unit="sat", @@ -233,21 +299,19 @@ async def test_recieve_token_uses_only_requested_destination_mint() -> None: proofs=[Mock(amount=100)], ) source_wallet = Mock() - swap = AsyncMock(return_value=(99, "sat", destination)) + redeem = AsyncMock(return_value=(99, "sat", source)) with ( - patch.object(settings, "primary_mint", destination), - patch.object(settings, "cashu_mints", [destination]), + patch.object(settings, "primary_mint", primary), + patch.object(settings, "cashu_mints", [primary, source]), patch("routstr.wallet.deserialize_token_from_string", return_value=token), patch("routstr.wallet.get_wallet", AsyncMock(return_value=source_wallet)), - patch("routstr.wallet.swap_to_trusted_mint", swap), + patch("routstr.wallet._redeem_same_mint", redeem), ): - result = await recieve_token( - "cashuAtoken", destination_mint=destination, destination_unit="sat" - ) + result = await recieve_token("cashuAtoken", destination_unit="sat") - assert result == (99, "sat", destination) - swap.assert_awaited_once_with(token, source_wallet, destination_mints=[destination]) + assert result == (99, "sat", source) + redeem.assert_awaited_once_with(source_wallet, token) @pytest.mark.asyncio @@ -256,35 +320,12 @@ async def test_recieve_token_rejects_unit_mismatch_before_wallet_mutation() -> N get_wallet = AsyncMock() with ( + patch("routstr.wallet.settings.cashu_mints", ["http://key-mint:3338"]), patch("routstr.wallet.deserialize_token_from_string", return_value=token), patch("routstr.wallet.get_wallet", get_wallet), pytest.raises(ValueError, match="liability unit"), ): - await recieve_token( - "cashuAtoken", - destination_mint="http://key-mint:3338", - destination_unit="sat", - ) - - get_wallet.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_recieve_token_cross_mint_output_unit_must_match() -> None: - token = Mock(mint="http://foreign:3338", unit="msat", keysets=["keyset"]) - get_wallet = AsyncMock() - - with ( - patch("routstr.wallet.deserialize_token_from_string", return_value=token), - patch("routstr.wallet.settings.primary_mint_unit", "sat"), - patch("routstr.wallet.get_wallet", get_wallet), - pytest.raises(ValueError, match="liability unit"), - ): - await recieve_token( - "cashuAtoken", - destination_mint="http://key-mint:3338", - destination_unit="msat", - ) + await recieve_token("cashuAtoken", destination_unit="sat") get_wallet.assert_not_awaited() @@ -392,6 +433,33 @@ async def test_send_token() -> None: assert token == "test_token" +@pytest.mark.asyncio +async def test_owner_only_token_rejects_customer_backed_proofs() -> None: + mint = "http://mint:3338" + proof = Mock(amount=1000, reserved=False) + wallet = Mock(keysets={}, proofs=[proof], select_to_send=AsyncMock()) + + with ( + patch( + "routstr.wallet.find_trusted_mint_with_funds", + AsyncMock(return_value=mint), + ), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + return_value=[proof], + ), + patch( + "routstr.wallet._owner_balance_for_mint_and_unit", + AsyncMock(return_value=50), + ), + pytest.raises(ValueError, match="Owner Cashu balance"), + ): + await send_token_from_owner_locked(100, "sat", mint) + + wallet.select_to_send.assert_not_awaited() + + @pytest.mark.asyncio async def test_release_token_reservation_unreserves_local_proofs() -> None: from routstr.wallet import release_token_reservation @@ -633,7 +701,41 @@ async def test_credit_balance() -> None: @pytest.mark.asyncio -async def test_credit_balance_constrains_redemption_to_key_mint() -> None: +async def test_concurrent_duplicate_token_credits_exactly_once() -> None: + key = Mock(balance=0, hashed_key="duplicate-key") + session = AsyncMock() + session.exec.return_value.rowcount = 1 + session.refresh = AsyncMock() + receive = AsyncMock( + side_effect=[ + (1000, "sat", "https://mint.test"), + ValueError("Mint Error: proofs already spent (Code: 11001)"), + ] + ) + store = AsyncMock() + + with ( + patch("routstr.wallet.recieve_token", receive), + patch("routstr.wallet.store_cashu_transaction", store), + ): + results = await asyncio.gather( + credit_balance("cashuAduplicate", key, session), + credit_balance("cashuAduplicate", key, session), + return_exceptions=True, + ) + + assert sum(result == 1_000_000 for result in results) == 1 + failure = next(result for result in results if isinstance(result, Exception)) + classified = classify_redemption_error(failure) + assert classified is not None and classified[3] == "cashu_token_already_spent" + assert session.exec.await_count == 1 + store.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_credit_balance_redeems_on_token_mint_not_key_mint() -> None: + """Top-ups are redeemed where the token was issued; the key's bound mint is + only a refund preference, never a swap destination.""" key_mint = "http://key-mint:3338" mock_key = Mock( balance=1_000_000, @@ -649,9 +751,7 @@ async def test_credit_balance_constrains_redemption_to_key_mint() -> None: with patch("routstr.wallet.store_cashu_transaction", AsyncMock()): await credit_balance("cashuAtoken", mock_key, mock_session) - receive.assert_awaited_once_with( - "cashuAtoken", destination_mint=key_mint, destination_unit="sat" - ) + receive.assert_awaited_once_with("cashuAtoken", destination_unit="sat") @pytest.mark.asyncio @@ -725,54 +825,6 @@ async def test_credit_balance_rejects_missing_key() -> None: assert not mock_session.commit.called -@pytest.mark.asyncio -async def test_swap_to_primary_mint_insufficient_for_fees() -> None: - """Token amount is less than melt_quote.amount + melt_quote.fee_reserve. - The quote mocks are static, so every retry observes the same shortfall — - the swap must still give up and raise.""" - from routstr.wallet import swap_to_primary_mint - - mock_token = Mock() - mock_token.mint = "http://foreign:3338" - mock_token.unit = "sat" - mock_token.amount = 404 - mock_token.keysets = ["keyset1"] - mock_token.proofs = [{"amount": 404}] - - mock_token_wallet = Mock() - mock_token_wallet.load_mint_keysets = AsyncMock() - mock_token_wallet.activate_keyset = AsyncMock() - mock_token_wallet._expand_short_keyset_ids = AsyncMock() - mock_token_wallet.load_proofs = AsyncMock() - mock_token_wallet.get_fees_for_proofs = Mock(return_value=0) - - mock_primary_wallet = Mock() - mock_primary_wallet.load_mint = AsyncMock() - mock_primary_wallet.load_proofs = AsyncMock() - - mock_mint_quote = Mock() - mock_mint_quote.quote = "mint_quote_123" - mock_mint_quote.request = "lnbc1..." - mock_primary_wallet.request_mint = AsyncMock(return_value=mock_mint_quote) - - mock_melt_quote = Mock() - mock_melt_quote.quote = "melt_quote_123" - mock_melt_quote.amount = 400 - mock_melt_quote.fee_reserve = 12 # total needed: 412 > 404 - mock_token_wallet.melt_quote = AsyncMock(return_value=mock_melt_quote) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(ValueError, match="insufficient to cover melt fees"): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - # melt should never have been called - mock_token_wallet.melt.assert_not_called() - - @pytest.mark.asyncio async def test_recieve_token_untrusted_mint() -> None: mock_wallet = Mock() @@ -785,724 +837,49 @@ async def test_recieve_token_untrusted_mint() -> None: mock_token.amount = 1000 mock_deserialize.return_value = mock_token - mock_wallet.load_mint = AsyncMock() - mock_wallet.load_proofs = AsyncMock() - with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): + with_db = AsyncMock(return_value=mock_wallet) + with patch("routstr.wallet.Wallet.with_db", with_db): + with pytest.raises(UntrustedSourceMintError): + await recieve_token("test_token") + with_db.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_recieve_token_accepts_multiple_keysets() -> None: + """TokenV4 groups proofs per keyset, and the node itself emits such tokens + after a mint rotates keysets, so redemption must span keysets.""" + mock_wallet = Mock() + mock_wallet.split = AsyncMock() + # Cashu sums input_fee_ppk per proof's own keyset. + mock_wallet.get_fees_for_proofs = Mock(return_value=2) + mock_wallet.load_mint_keysets = AsyncMock() + mock_wallet.activate_keyset = AsyncMock() + mock_wallet._expand_short_keyset_ids = AsyncMock() + + from routstr.core.settings import settings + + with patch.object(settings, "cashu_mints", ["http://mint:3338"]): + with patch("routstr.wallet.deserialize_token_from_string") as mock_deserialize: + mock_token = Mock() + mock_token.keysets = ["keyset1", "keyset2"] + mock_token.mint = "http://mint:3338" + mock_token.unit = "sat" + mock_token.amount = 1000 + mock_token.proofs = [ + {"amount": 600, "id": "keyset1"}, + {"amount": 400, "id": "keyset2"}, + ] + mock_deserialize.return_value = mock_token + with patch( - "routstr.wallet.swap_to_trusted_mint", - return_value=(900, "sat", "http://mint:3338"), + "routstr.wallet.get_wallet", + AsyncMock(return_value=mock_wallet), ): - amount, unit, mint = await recieve_token("test_token") - assert amount == 900 - assert unit == "sat" - assert mint == "http://mint:3338" + amount, unit, mint = await recieve_token("cashuAmultikeyset") - -@pytest.mark.asyncio -async def test_swap_to_primary_mint_already_on_primary() -> None: - """Same-mint shortcut: the token is already on the primary mint. - - No cross-mint swap (no melt/mint), but the same-mint split(include_fees=True) - still burns the mint's NUT-02 input fee, so the credited amount must be face - minus the input fee — not full face value (the over-credit bug). DLEQ is - verified too, matching the trusted same-mint receive path. - """ - from routstr.core.settings import settings - from routstr.wallet import swap_to_primary_mint - - mock_token = Mock() - mock_token.mint = settings.primary_mint - mock_token.keysets = ["keyset1"] - mock_token.amount = 1000 - mock_token.unit = "sat" - mock_token.proofs = [{"amount": 1000}] - - mock_token_wallet = Mock() - mock_token_wallet.load_mint_keysets = AsyncMock() - mock_token_wallet.activate_keyset = AsyncMock() - mock_token_wallet._expand_short_keyset_ids = AsyncMock() - mock_token_wallet.load_proofs = AsyncMock() - mock_token_wallet.verify_proofs_dleq = Mock() - # Mock a 3-sat input fee from the Cashu wallet API. - mock_token_wallet.get_fees_for_proofs = Mock(return_value=3) - mock_token_wallet.split = AsyncMock(return_value=None) - mock_token_wallet.request_mint = AsyncMock() - mock_token_wallet.melt_quote = AsyncMock() - - with patch("routstr.wallet.get_wallet", AsyncMock(return_value=mock_token_wallet)): - amount, unit, mint = await swap_to_primary_mint(mock_token, mock_token_wallet) - - assert amount == 997 # 1000 face - 3 sat input fee - assert unit == "sat" - assert mint == settings.primary_mint - mock_token_wallet.verify_proofs_dleq.assert_called_once_with(mock_token.proofs) - mock_token_wallet.get_fees_for_proofs.assert_called_once_with(mock_token.proofs) - mock_token_wallet.split.assert_called_once() - mock_token_wallet.request_mint.assert_not_called() - mock_token_wallet.melt_quote.assert_not_called() - - -# --------------------------------------------------------------------------- -# Swap fee estimation and reactive retry -# -# Spec: the estimation pass subtracts only observed fees (no safety buffer). -# swap_to_primary_mint then runs the mint-quote/melt-quote/melt cycle and, when -# the foreign mint demands more than estimated (at quote or at melt time), -# retries with the amount recomputed from the observed fee — at most 3 attempts. -# Melt failures unrelated to fees are not retried. -# --------------------------------------------------------------------------- - - -def _make_swap_mocks( - token_amount: int, - fee_reserves: list[int], - input_fees: int = 0, - mint_url: str = "http://foreign-mint:3338", -) -> tuple[Mock, Mock, Mock]: - """Return (token, token_wallet, primary_wallet) mocks that act like a mint. - - Mint quotes pass the requested amount through their ``request`` field and - melt quotes echo that amount back, so the mocks stay consistent for - whatever amounts the implementation requests. ``fee_reserves`` supplies the - fee_reserve of each successive melt quote (the first serves the estimation - pass); requesting more quotes than provided fails the test. - """ - mock_token = Mock() - mock_token.mint = mint_url - mock_token.unit = "sat" - mock_token.amount = token_amount - mock_token.keysets = ["keyset1"] - mock_token.proofs = [Mock(amount=token_amount)] - - mock_token_wallet = Mock() - mock_token_wallet.load_mint_keysets = AsyncMock() - mock_token_wallet.activate_keyset = AsyncMock() - mock_token_wallet._expand_short_keyset_ids = AsyncMock() - mock_token_wallet.load_proofs = AsyncMock() - mock_token_wallet.get_fees_for_proofs = Mock(return_value=input_fees) - - mock_primary_wallet = Mock() - mock_primary_wallet.load_mint = AsyncMock() - mock_primary_wallet.load_proofs = AsyncMock() - mock_primary_wallet.available_balance = Mock(amount=0) - mock_primary_wallet.mint = AsyncMock(return_value=Mock()) - - fees = iter(fee_reserves) - - def _next_fee() -> int: - try: - return next(fees) - except StopIteration: - raise AssertionError( - "more melt quotes requested than fee_reserves provided" - ) from None - - mock_primary_wallet.request_mint = AsyncMock( - side_effect=lambda amount: Mock(quote=f"mint_quote_{amount}", request=amount) - ) - mock_token_wallet.melt_quote = AsyncMock( - side_effect=lambda invoice: Mock( - quote=f"melt_quote_{invoice}", amount=invoice, fee_reserve=_next_fee() - ) - ) - mock_token_wallet.melt = AsyncMock(return_value=Mock(state=MeltQuoteState.paid)) - - return mock_token, mock_token_wallet, mock_primary_wallet - - -@pytest.mark.asyncio -async def test_swap_to_primary_mint_success() -> None: - """No retry needed: real quote matches the estimate, full net amount minted.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - amount, unit, mint = await swap_to_primary_mint( - mock_token, mock_token_wallet - ) - - assert amount == 990 # 1000 - fee_reserve(10), no buffer subtracted - assert unit == "sat" - assert mint == "http://primary:3338" - assert mock_primary_wallet.request_mint.call_count == 2 - mock_primary_wallet.request_mint.assert_any_call(1000) - mock_primary_wallet.request_mint.assert_any_call(990) - assert mock_token_wallet.melt_quote.call_count == 2 - assert mock_token_wallet.melt.call_count == 1 - assert mock_primary_wallet.mint.called - - -@pytest.mark.asyncio -@pytest.mark.parametrize("fee_reserve", [1, 10, 100]) -async def test_calculate_swap_amount_subtracts_only_observed_fees( - fee_reserve: int, -) -> None: - """Estimation: minted_amount = token - fee_reserve, with no safety buffer.""" - from routstr.wallet import _calculate_swap_amount - - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[fee_reserve] - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - result = await _calculate_swap_amount( - amount_msat=1_000_000, - token_unit="sat", - token_mint_url="http://foreign-mint:3338", - token_wallet=mock_token_wallet, - primary_wallet=mock_primary_wallet, - proofs=[], - ) - - assert result == 1000 - fee_reserve - - -@pytest.mark.asyncio -async def test_calculate_swap_amount_includes_input_fees() -> None: - """Estimation subtracts NUT-02 input fees alongside the melt fee_reserve.""" - from routstr.wallet import _calculate_swap_amount - - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 500, fee_reserves=[10], input_fees=3 - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - result = await _calculate_swap_amount( - amount_msat=500_000, - token_unit="sat", - token_mint_url="http://foreign-mint:3338", - token_wallet=mock_token_wallet, - primary_wallet=mock_primary_wallet, - proofs=[], - ) - - assert result == 487 # 500 - 10 - 3 - - -@pytest.mark.asyncio -async def test_swap_retries_when_real_quote_exceeds_estimate() -> None: - """The real melt quote demands a higher fee than the estimate (20 → 23). - Instead of failing, the swap recomputes the amount from the observed fee - and re-quotes: 1000 - 23 = 977, which fits (977 + 23 <= 1000).""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[20, 23, 23] - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - amount, unit, mint = await swap_to_primary_mint( - mock_token, mock_token_wallet - ) - - assert amount == 977 - assert unit == "sat" - mock_primary_wallet.request_mint.assert_any_call(980) - mock_primary_wallet.request_mint.assert_any_call(977) - assert mock_token_wallet.melt_quote.call_count == 3 # estimation + 2 attempts - assert mock_token_wallet.melt.call_count == 1 - - -@pytest.mark.asyncio -async def test_swap_retries_when_melt_demands_more_than_quoted() -> None: - """The mint.cubabitcoin.org incident: every quote reports fee_reserve=1, - but the mint demands 2 sats at melt time ("Provided: 179, needed: 180"). - The swap must retry with a smaller invoice (177) so the second melt fits, - instead of failing the topup.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 179, fee_reserves=[1, 1, 1], mint_url="http://mint.cubabitcoin.org" - ) - mock_token_wallet.melt.side_effect = [ - Exception( - "Mint Error: not enough inputs provided for melt. " - "Provided: 179, needed: 180 (Code: 11000)" - ), - Mock(state=MeltQuoteState.paid), - ] - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - amount, unit, mint = await swap_to_primary_mint( - mock_token, mock_token_wallet - ) - - assert amount == 177 # 179 - 1 (estimate) - 1 (observed melt shortfall) - assert mock_token_wallet.melt.call_count == 2 - mock_primary_wallet.request_mint.assert_any_call(178) - mock_primary_wallet.request_mint.assert_any_call(177) - - -@pytest.mark.asyncio -async def test_swap_retries_on_cdk_unbalanced_error() -> None: - """cdk-based mints report insufficient melt inputs as the registered code - 11005 (TransactionUnbalanced) with their own message wording — no - Provided/needed amounts to parse. The retry must classify it by code and - fall back to shrinking by 1.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 179, fee_reserves=[1, 1, 1] - ) - mock_token_wallet.melt.side_effect = [ - Exception("Mint Error: Transaction unbalanced: 179, 178, 2 (Code: 11005)"), - Mock(state=MeltQuoteState.paid), - ] - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - amount, unit, mint = await swap_to_primary_mint( - mock_token, mock_token_wallet - ) - - assert amount == 177 - assert mock_token_wallet.melt.call_count == 2 - - -@pytest.mark.asyncio -async def test_swap_quote_retries_exhausted() -> None: - """A mint that escalates fee_reserve on every re-quote exhausts the retry - budget (3 attempts) and fails cleanly; melt is never executed.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[1, 10, 25, 50] - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(ValueError, match="insufficient to cover melt fees"): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - assert mock_token_wallet.melt_quote.call_count == 4 # estimation + 3 attempts - mock_token_wallet.melt.assert_not_called() - - -@pytest.mark.asyncio -async def test_swap_melt_retries_exhausted() -> None: - """A mint that always demands more at melt time than it quoted exhausts - the retry budget; the last melt failure is wrapped as ValueError.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 5000, fee_reserves=[50, 50, 50, 50] - ) - mock_token_wallet.melt = AsyncMock( - side_effect=Exception( - "Mint Error: not enough inputs provided for melt. " - "Provided: 5000, needed: 5200 (Code: 11000)" - ) - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(ValueError, match="Failed to melt token"): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - assert mock_token_wallet.melt.call_count == 3 - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "primary_unit,token_unit,amount_msat,fees,expected", - [ - ("sat", "sat", 179_000, 2, 177), - ("msat", "sat", 179_000, 2, 177_000), - ("sat", "msat", 179_000, 2_000, 177), - ], -) -async def test_net_minted_amount_unit_conversions( - primary_unit: str, token_unit: str, amount_msat: int, fees: int, expected: int -) -> None: - """Fee subtraction converts correctly between sat and msat on either side.""" - from routstr.core.settings import settings - from routstr.wallet import _net_minted_amount - - with patch.object(settings, "primary_mint_unit", primary_unit): - assert _net_minted_amount(amount_msat, token_unit, fees) == expected - - -@pytest.mark.parametrize( - "message,expected", - [ - # nutshell: retryable with exact shortfall from the detail text - ( - "Mint Error: not enough inputs provided for melt. " - "Provided: 179, needed: 182 (Code: 11000)", - 3, - ), - # verbatim production error from issue #468, including cashu-py's - # "could not pay invoice" wrapper around the mint detail - ( - "could not pay invoice: Mint Error: not enough inputs provided " - "for melt. Provided: 179, needed: 180 (Code: 11000)", - 1, - ), - # cdk: registered TransactionUnbalanced code, no parsable amounts - ("Mint Error: Transaction unbalanced: 179, 178, 2 (Code: 11005)", 1), - # nutshell wording without a code suffix - ("not enough inputs provided for melt", 1), - # nonsensical amounts (needed <= provided) fall back to the minimal step - ( - "Mint Error: not enough inputs provided for melt. " - "Provided: 180, needed: 179 (Code: 11000)", - 1, - ), - # a generic 11000 without the shortfall text is not a fee shortfall: - # 11000 is nutshell's catch-all TransactionError, so retrying (shrinking - # the invoice) would never help and only masks the real error - ("Mint Error: Duplicate inputs provided. (Code: 11000)", None), - # spent proofs must never be retried: the funds are gone - ("Mint Error: Token already spent. (Code: 11001)", None), - # Lightning failures must never be retried: a smaller invoice won't help - ("Mint Error: Lightning payment failed. (Code: 20004)", None), - # unrecognizable errors (timeouts, bugs) must never be retried - ("Connection timeout", None), - ], -) -def test_melt_shortfall_classifier(message: str, expected: int | None) -> None: - """Retry classification across mint implementations and failure classes.""" - from routstr.wallet import _melt_insufficient_shortfall - - assert _melt_insufficient_shortfall(Exception(message)) == expected - - -@pytest.mark.asyncio -async def test_calculate_swap_amount_same_mint_short_circuit() -> None: - """When the token is already on the primary mint no fees apply and no - quotes are requested.""" - from routstr.wallet import _calculate_swap_amount - - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(1000, fee_reserves=[]) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - result = await _calculate_swap_amount( - amount_msat=1_000_000, - token_unit="sat", - token_mint_url="http://primary:3338", - token_wallet=mock_token_wallet, - primary_wallet=mock_primary_wallet, - proofs=[], - ) - - assert result == 1000 - mock_primary_wallet.request_mint.assert_not_called() - mock_token_wallet.melt_quote.assert_not_called() - - -@pytest.mark.asyncio -async def test_calculate_swap_amount_msat_primary_unit() -> None: - """With an msat primary mint the dummy quote and result stay in msats.""" - from routstr.wallet import _calculate_swap_amount - - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(179, fee_reserves=[2]) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "msat"): - result = await _calculate_swap_amount( - amount_msat=179_000, - token_unit="sat", - token_mint_url="http://foreign-mint:3338", - token_wallet=mock_token_wallet, - primary_wallet=mock_primary_wallet, - proofs=[], - ) - - assert result == 177_000 # 179_000 msat - 2 sat fee - mock_primary_wallet.request_mint.assert_called_once_with(179_000) - - -@pytest.mark.asyncio -async def test_calculate_swap_amount_fees_exceed_token() -> None: - """Fees larger than the token itself fail fast, before any melt.""" - from routstr.wallet import _calculate_swap_amount - - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 179, fee_reserves=[200] - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with pytest.raises(ValueError, match="exceed token amount"): - await _calculate_swap_amount( - amount_msat=179_000, - token_unit="sat", - token_mint_url="http://foreign-mint:3338", - token_wallet=mock_token_wallet, - primary_wallet=mock_primary_wallet, - proofs=[], - ) - - -@pytest.mark.asyncio -async def test_calculate_swap_amount_wraps_estimation_failure() -> None: - """Estimation infrastructure failures surface as a single clear ValueError.""" - from routstr.wallet import _calculate_swap_amount - - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(179, fee_reserves=[]) - mock_primary_wallet.request_mint = AsyncMock(side_effect=Exception("mint offline")) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with pytest.raises(ValueError, match="Failed to estimate fees"): - await _calculate_swap_amount( - amount_msat=179_000, - token_unit="sat", - token_mint_url="http://foreign-mint:3338", - token_wallet=mock_token_wallet, - primary_wallet=mock_primary_wallet, - proofs=[], - ) - - -@pytest.mark.asyncio -async def test_swap_coerces_non_integer_amount() -> None: - """Token amounts arriving as floats are coerced before any arithmetic.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - mock_token.amount = 1000.0 - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - amount, unit, mint = await swap_to_primary_mint( - mock_token, mock_token_wallet - ) - - assert amount == 990 - assert isinstance(amount, int) - - -@pytest.mark.asyncio -async def test_swap_rejects_unknown_unit() -> None: - """Units other than sat/msat are rejected before any quote is requested.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[] - ) - mock_token.unit = "usd" - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(ValueError, match="Invalid unit"): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - mock_primary_wallet.request_mint.assert_not_called() - - -@pytest.mark.asyncio -async def test_swap_msat_token_already_on_primary() -> None: - """msat-denominated tokens on the primary mint short-circuit unchanged.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, _ = _make_swap_mocks( - 179_000, fee_reserves=[], mint_url="http://primary:3338" - ) - mock_token.unit = "msat" - mock_token_wallet.split = AsyncMock() - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_token_wallet): - amount, unit, mint = await swap_to_primary_mint( - mock_token, mock_token_wallet - ) - - assert (amount, unit, mint) == (179_000, "msat", "http://primary:3338") - - -# --------------------------------------------------------------------------- -# Mint-on-primary failure handling after a successful melt -# -# At this point the foreign proofs are already spent: failures here mean funds -# are in limbo, so errors must propagate (never be swallowed) and recovery must -# never credit proofs the wallet does not actually hold. -# --------------------------------------------------------------------------- - - -def _with_recovery_mocks( - mock_primary_wallet: Mock, mint_error: str, balances: list[int] -) -> None: - """Make primary mint() fail and stage available_balance per load_proofs call.""" - mock_primary_wallet.mint = AsyncMock(side_effect=Exception(mint_error)) - mock_primary_wallet.keysets = ["keyset_primary"] - balance_iter = iter(balances) - - def advance_balance(reload: bool = False) -> None: - mock_primary_wallet.available_balance = Mock(amount=next(balance_iter)) - - mock_primary_wallet.load_proofs = AsyncMock(side_effect=advance_balance) - mock_primary_wallet.restore_tokens_for_keyset = AsyncMock() - - -@pytest.mark.asyncio -async def test_swap_mint_failure_after_melt_is_token_consumed() -> None: - """A non-recoverable mint failure after melt is a non-retryable - TokenConsumedError (the melt already spent the foreign proofs), with the - original error preserved in the cause chain.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - _with_recovery_mocks( - mock_primary_wallet, "Mint Error: Quote is expired (Code: 20007)", [0] - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(TokenConsumedError) as exc_info: - await swap_to_primary_mint(mock_token, mock_token_wallet) - - assert "Quote is expired" in str(exc_info.value.__cause__) - assert mock_token_wallet.melt.call_count == 1 - mock_primary_wallet.restore_tokens_for_keyset.assert_not_called() - - -@pytest.mark.asyncio -async def test_swap_recovers_orphaned_proofs_on_outputs_already_signed() -> None: - """11003 (outputs already signed): a recovery scan that restores the full - minted amount lets the swap complete normally.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - _with_recovery_mocks( - mock_primary_wallet, - "Mint Error: outputs already signed (Code: 11003)", - [0, 990], # pre-mint balance, post-recovery balance - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - amount, unit, mint = await swap_to_primary_mint( - mock_token, mock_token_wallet - ) - - assert amount == 990 - mock_primary_wallet.restore_tokens_for_keyset.assert_awaited_once_with( - "keyset_primary", to=1, batch=25 - ) - - -@pytest.mark.asyncio -async def test_swap_recovery_shortfall_refuses_credit() -> None: - """When the recovery scan restores less than the minted amount, the swap - must fail rather than credit proofs the wallet does not hold.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - _with_recovery_mocks( - mock_primary_wallet, - "Mint Error: outputs already signed (Code: 11003)", - [0, 100], # recovery restores only 100 of the expected 990 - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(TokenConsumedError, match="Swap recovery failed"): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - -@pytest.mark.asyncio -async def test_swap_recovery_failure_wrapped() -> None: - """When the recovery scan itself fails, the error is wrapped and raised — - never swallowed.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - _with_recovery_mocks( - mock_primary_wallet, - "Mint Error: outputs already signed (Code: 11003)", - [0], - ) - mock_primary_wallet.restore_tokens_for_keyset = AsyncMock( - side_effect=Exception("wallet db locked") - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(TokenConsumedError, match="recovery unsuccessful"): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - -@pytest.mark.asyncio -async def test_recieve_token_rejects_multiple_keysets() -> None: - """Multi-keyset tokens are rejected before touching any wallet.""" - with patch("routstr.wallet.deserialize_token_from_string") as mock_deserialize: - mock_token = Mock() - mock_token.keysets = ["keyset1", "keyset2"] - mock_deserialize.return_value = mock_token - - with pytest.raises(ValueError, match="Multiple keysets"): - await recieve_token("cashuAmultikeyset") + assert (amount, unit, mint) == (998, "sat", "http://mint:3338") + mock_wallet.get_fees_for_proofs.assert_called_once_with(mock_token.proofs) + mock_wallet.split.assert_awaited_once() @pytest.mark.asyncio @@ -1552,31 +929,6 @@ async def test_credit_balance_propagates_audit_store_failure_after_credit() -> N assert mock_session.commit.called -@pytest.mark.asyncio -async def test_swap_does_not_retry_on_payment_failure() -> None: - """Melt failures unrelated to fees (e.g. routing failure) are not retried: - a smaller invoice would not help, and the error must surface immediately.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - mock_token_wallet.melt = AsyncMock( - side_effect=Exception("Mint Error: Lightning payment failed. (Code: 20004)") - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(ValueError, match="Failed to melt token"): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - assert mock_token_wallet.melt.call_count == 1 - assert mock_primary_wallet.request_mint.call_count == 2 - - # --- Mint-unreachable classification (is_mint_connection_error) --------------- @@ -1597,7 +949,7 @@ def test_rate_limited_mint_is_classified_as_unreachable() -> None: assert classify_redemption_error(error) == ( "mint_rate_limited", 503, - "Cashu mint rate-limited; retry after cooldown", + "Cashu mint is rate-limiting requests; retry later", "cashu_mint_rate_limited", ) @@ -1707,38 +1059,6 @@ def test_classify_generic_valueerror_is_not_zero_value() -> None: ) -@pytest.mark.asyncio -async def test_swap_mint_transport_error_after_melt_is_not_retryable() -> None: - """A transport error minting on the primary mint (after the foreign melt - already spent the proofs) classifies as a non-retryable token_consumed 500, - never a retryable mint_unreachable 503.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - # Melt succeeds (proofs spent); minting on primary hits a transport error. - mock_primary_wallet.mint = AsyncMock( - side_effect=httpx.ConnectError("primary mint down") - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(TokenConsumedError) as exc_info: - await swap_to_primary_mint(mock_token, mock_token_wallet) - - classified = classify_redemption_error(exc_info.value) - assert classified is not None - _type, status, _msg, code = classified - assert status == 500 - assert code == "cashu_token_consumed" - assert is_mint_connection_error(exc_info.value) is False - assert mock_token_wallet.melt.call_count == 1 - - @pytest.mark.asyncio async def test_credit_balance_db_transport_error_is_token_consumed() -> None: """A transport-like DB failure after the token is redeemed must be @@ -1759,64 +1079,6 @@ async def test_credit_balance_db_transport_error_is_token_consumed() -> None: assert is_mint_connection_error(exc_info.value) is False -@pytest.mark.asyncio -async def test_swap_fee_estimation_transport_error_raises_mint_connection_error() -> ( - None -): - """A transport failure while estimating fees is surfaced as - MintConnectionError (→ 503), not a generic fee ValueError (→ 422).""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10] - ) - mock_primary_wallet.request_mint = AsyncMock( - side_effect=httpx.ConnectError("All connection attempts failed") - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(MintConnectionError): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - mock_token_wallet.melt.assert_not_called() - - -@pytest.mark.asyncio -async def test_swap_melt_transport_error_is_never_reported_reusable() -> None: - """A timed-out melt remains ambiguous even when an immediate snapshot says - UNPAID/UNSPENT, so callers must not receive the original token for retry.""" - from routstr.wallet import swap_to_primary_mint - - mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[10, 10] - ) - mock_token_wallet.melt = AsyncMock(side_effect=httpx.ConnectTimeout("timed out")) - from cashu.core.base import MeltQuoteState, ProofSpentState - - mock_token_wallet.get_melt_quote = AsyncMock( - return_value=Mock(state=MeltQuoteState.unpaid) - ) - mock_token_wallet.check_proof_state = AsyncMock( - return_value=Mock( - states=[Mock(state=ProofSpentState.unspent) for _ in mock_token.proofs] - ) - ) - - from routstr.core.settings import settings - - with patch.object(settings, "primary_mint", "http://primary:3338"): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(TokenConsumedError, match="ambiguous"): - await swap_to_primary_mint(mock_token, mock_token_wallet) - - assert mock_token_wallet.melt.call_count == 1 - - @pytest.mark.asyncio async def test_execute_bolt11_payment_rejects_unpaid_melt_state() -> None: plan = MagicMock() @@ -2428,78 +1690,6 @@ async def test_lightning_mint_fallback_for_topups() -> None: mock_secondary_wallet.request_mint.assert_called_once() -@pytest.mark.asyncio -async def test_swap_falls_back_when_primary_wallet_cannot_load() -> None: - from routstr.core.settings import settings - from routstr.wallet import swap_to_primary_mint - - primary = "http://primary:3338" - secondary = "http://secondary:3338" - foreign = "http://foreign:3338" - - token = Mock( - mint=foreign, - unit="sat", - amount=1000, - keysets=["keyset1"], - proofs=[Mock(amount=1000)], - ) - source_wallet = Mock( - load_mint_keysets=AsyncMock(), - activate_keyset=AsyncMock(), - _expand_short_keyset_ids=AsyncMock(), - load_proofs=AsyncMock(), - get_fees_for_proofs=Mock(return_value=0), - melt_quote=AsyncMock( - return_value=Mock(quote="melt_q", amount=990, fee_reserve=10) - ), - melt=AsyncMock(return_value=Mock(state=MeltQuoteState.paid)), - ) - - mint_quote = Mock(quote="mint_q_secondary", request="lnbc1secondary") - secondary_wallet = Mock( - load_mint=AsyncMock(), - load_proofs=AsyncMock(), - available_balance=Mock(amount=0), - keysets=["ks_secondary"], - restore_tokens_for_keyset=AsyncMock(), - request_mint=AsyncMock(return_value=mint_quote), - mint=AsyncMock(return_value=Mock()), - ) - - async def get_wallet(mint: str, *args: object, **kwargs: object) -> Mock: - if mint == primary: - raise httpx.ConnectError("primary down") - return secondary_wallet - - mock_get = AsyncMock(side_effect=get_wallet) - with ( - patch.object(settings, "primary_mint", primary), - patch.object(settings, "primary_mint_unit", "sat"), - patch.object(settings, "cashu_mints", [primary, secondary]), - patch.object(settings, "mint_max_concurrency", 0), - patch.object(settings, "mint_operation_timeout_seconds", 0), - patch("asyncio.sleep", AsyncMock()), - patch("routstr.wallet.get_wallet", side_effect=mock_get), - patch("routstr.wallet.logger.warning") as warning, - patch("routstr.wallet.logger.info") as info, - ): - amount, unit, mint_url = await swap_to_primary_mint(token, source_wallet) - - assert (amount, unit, mint_url) == (990, "sat", secondary) - secondary_wallet.mint.assert_awaited_once() - assert mock_get.await_args_list[0].args[0] == primary - assert any(call.args[0] == secondary for call in mock_get.await_args_list) - events = { - call.kwargs["extra"]["event"] - for call in [*warning.call_args_list, *info.call_args_list] - if "extra" in call.kwargs and "event" in call.kwargs["extra"] - } - assert "cashu_destination_failed" in events - assert "cashu_destination_selected" in events - assert "cashu_swap_completed" in events - - def test_raise_on_error_request_identifies_cashu_mint_error() -> None: from routstr.mint import MintError from routstr.wallet import Wallet @@ -2610,139 +1800,6 @@ async def test_lightning_mint_fallback_rejects_zero_amount() -> None: await _request_mint_with_fallback(-5) -@pytest.mark.asyncio -async def test_wallet_request_mint_fallback_rejects_zero_amount() -> None: - """Zero or negative amounts must be rejected before reaching the mint.""" - from routstr.wallet import _request_mint_with_fallback - - with pytest.raises(ValueError, match="amount must be > 0"): - await _request_mint_with_fallback(0, op_name="test") - - with pytest.raises(ValueError, match="amount must be > 0"): - await _request_mint_with_fallback(-1, op_name="test") - - -@pytest.mark.asyncio -async def test_wallet_fallback_on_429_no_in_place_retry() -> None: - """A 429 from the primary mint must trigger immediate fallback to the - secondary — _mint_operation must NOT retry in-place when - retry_on_rate_limit=False is set by _request_mint_with_fallback.""" - from routstr.core.settings import settings - from routstr.wallet import _request_mint_with_fallback - - primary = "http://primary:3338" - secondary = "http://secondary:3338" - - request = httpx.Request("POST", "http://primary:3338/v1/mint/quote/bolt11") - response = httpx.Response(429, request=request, headers={"Retry-After": "60"}) - primary_call_count = 0 - - async def primary_request_mint(_amount: int) -> None: - nonlocal primary_call_count - primary_call_count += 1 - raise httpx.HTTPStatusError("rate limited", request=request, response=response) - - mock_primary_wallet = Mock() - mock_primary_wallet.request_mint = AsyncMock(side_effect=primary_request_mint) - - mock_quote = Mock(quote="q_secondary", request="lnbc1secondary") - mock_secondary_wallet = Mock() - mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote) - - wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} - mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) - - with patch.object(settings, "primary_mint", primary): - with patch.object(settings, "cashu_mints", [primary, secondary]): - with patch.object(settings, "mint_retry_max_attempts", 3): - with patch.object(settings, "mint_max_concurrency", 0): - with patch.object(settings, "mint_operation_timeout_seconds", 0): - with patch("asyncio.sleep", AsyncMock()) as mock_sleep: - with patch( - "routstr.wallet.get_wallet", side_effect=mock_get - ): - _, mint_url, _ = await _request_mint_with_fallback( - 1000, op_name="test_429_fallback" - ) - - assert mint_url == secondary - assert primary_call_count == 1 - mock_secondary_wallet.request_mint.assert_called_once() - mock_sleep.assert_not_called() - - -@pytest.mark.asyncio -async def test_wallet_fallback_on_timeout_no_in_place_retry() -> None: - """A timeout from one destination must immediately try the next mint.""" - from routstr.core.settings import settings - from routstr.wallet import _request_mint_with_fallback - - primary = "http://primary:3338" - secondary = "http://secondary:3338" - primary_wallet = Mock( - request_mint=AsyncMock(side_effect=httpx.TimeoutException("timed out")) - ) - quote = Mock(quote="q_secondary", request="lnbc1secondary") - secondary_wallet = Mock(request_mint=AsyncMock(return_value=quote)) - wallets = {primary: primary_wallet, secondary: secondary_wallet} - - with ( - patch.object(settings, "primary_mint", primary), - patch.object(settings, "cashu_mints", [primary, secondary]), - patch.object(settings, "mint_retry_max_attempts", 3), - patch.object(settings, "mint_max_concurrency", 0), - patch.object(settings, "mint_operation_timeout_seconds", 0), - patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep, - patch( - "routstr.wallet.get_wallet", - AsyncMock(side_effect=lambda mint, *args, **kwargs: wallets[mint]), - ), - ): - _, mint_url, _ = await _request_mint_with_fallback( - 1000, op_name="test_timeout_fallback" - ) - - assert mint_url == secondary - primary_wallet.request_mint.assert_awaited_once_with(1000) - secondary_wallet.request_mint.assert_awaited_once_with(1000) - sleep.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_wallet_fallback_skips_mint_during_cooldown() -> None: - from routstr.core.settings import settings - from routstr.wallet import _MintRateGuard, _request_mint_with_fallback - - primary = "http://primary:3338" - secondary = "http://secondary:3338" - primary_wallet = Mock(request_mint=AsyncMock()) - quote = Mock(quote="q_secondary", request="lnbc1secondary") - secondary_wallet = Mock(request_mint=AsyncMock(return_value=quote)) - wallets = {primary: primary_wallet, secondary: secondary_wallet} - - with ( - patch.object(settings, "primary_mint", primary), - patch.object(settings, "cashu_mints", [primary, secondary]), - patch.object(settings, "mint_max_concurrency", 0), - patch.object(settings, "mint_operation_timeout_seconds", 0), - patch("routstr.mint.time.monotonic", return_value=10), - patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep, - patch( - "routstr.wallet.get_wallet", - AsyncMock(side_effect=lambda mint, *args, **kwargs: wallets[mint]), - ), - ): - _MintRateGuard.get(primary).apply_cooldown(60) - _, mint_url, _ = await _request_mint_with_fallback( - 1000, op_name="test_cooldown_fallback" - ) - - assert mint_url == secondary - primary_wallet.request_mint.assert_not_awaited() - secondary_wallet.request_mint.assert_awaited_once_with(1000) - sleep.assert_not_awaited() - - @pytest.mark.asyncio async def test_lightning_fallback_on_429_no_in_place_retry() -> None: """Same as above but for the lightning.py _request_mint_with_fallback.""" @@ -2970,6 +2027,7 @@ async def test_load_mint_propagates_rate_limit() -> None: from routstr.wallet import Wallet wallet = Wallet.__new__(Wallet) + wallet.url = "https://rate-limited-mint.example" error = MintRateLimitedError( "Cashu mint rate limited", request=httpx.Request("GET", "https://mint.example/v1/keysets"), @@ -2987,6 +2045,7 @@ async def test_load_mint_propagates_connection_error() -> None: from routstr.wallet import Wallet wallet = Wallet.__new__(Wallet) + wallet.url = "https://unavailable-mint.example" error = httpx.ConnectError("mint unavailable") with ( patch.object(wallet, "load_mint_keysets", new=AsyncMock(side_effect=error)), @@ -3002,6 +2061,7 @@ async def test_load_mint_runs_keysets_activation_and_info() -> None: from routstr.wallet import Wallet wallet = Wallet.__new__(Wallet) + wallet.url = "https://mint-load.example" with ( patch.object(wallet, "load_mint_keysets", new=AsyncMock()) as load_keysets, patch.object(wallet, "activate_keyset", new=AsyncMock()) as activate, diff --git a/tests/unit/test_x_cashu_json_body_with_data_prefix.py b/tests/unit/test_x_cashu_json_body_with_data_prefix.py new file mode 100644 index 00000000..61dfcf05 --- /dev/null +++ b/tests/unit/test_x_cashu_json_body_with_data_prefix.py @@ -0,0 +1,174 @@ +import json +import os +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.payment.cost_calculation import CostData # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider, _is_sse_body # noqa: E402 + +REFUND_TOKEN = "cashuBrefundtoken0123456789" + + +def _cost_data() -> CostData: + return CostData( + base_msats=0, + input_msats=2500, + output_msats=1500, + total_msats=4000, + total_usd=0.0002, + input_tokens=12, + output_tokens=8, + ) + + +def _chat_json_with_data_prefix() -> dict[str, Any]: + return { + "id": "chatcmpl-1", + "object": "chat.completion", + "model": "gpt-5-mini", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": { + "role": "assistant", + "content": "Here you go: data:image/png;base64,iVBORw0KGgo=", + }, + } + ], + "usage": {"prompt_tokens": 12, "completion_tokens": 8, "total_tokens": 20}, + } + + +def _responses_json_with_data_prefix() -> dict[str, Any]: + return { + "id": "resp-1", + "object": "response", + "model": "gpt-5-mini", + "output": [ + { + "type": "message", + "content": [{"type": "output_text", "text": "use data: prefix"}], + } + ], + "usage": {"input_tokens": 12, "output_tokens": 8, "total_tokens": 20}, + } + + +def _json_response(payload: dict[str, Any], content_type: str | None) -> httpx.Response: + headers = {"content-type": content_type} if content_type else {} + return httpx.Response(200, headers=headers, content=json.dumps(payload).encode()) + + +async def _settle_chat(response: httpx.Response) -> tuple[Any, AsyncMock]: + provider = BaseUpstreamProvider(base_url="http://test", api_key="test-key") + send_refund = AsyncMock(return_value=REFUND_TOKEN) + with ( + patch.object( + provider, "get_x_cashu_cost", new=AsyncMock(return_value=_cost_data()) + ), + patch.object(provider, "send_refund", new=send_refund), + ): + result = await provider.handle_x_cashu_chat_completion( + response=response, + amount=10_000, + unit="msat", + max_cost_for_model=9_000, + mint=None, + ) + return result, send_refund + + +async def _settle_responses(response: httpx.Response) -> tuple[Any, AsyncMock]: + provider = BaseUpstreamProvider(base_url="http://test", api_key="test-key") + send_refund = AsyncMock(return_value=REFUND_TOKEN) + with ( + patch.object( + provider, "get_x_cashu_cost", new=AsyncMock(return_value=_cost_data()) + ), + patch.object(provider, "send_refund", new=send_refund), + ): + result = await provider.handle_x_cashu_responses_completion( + response=response, + amount=10_000, + unit="msat", + max_cost_for_model=9_000, + mint=None, + ) + return result, send_refund + + +@pytest.mark.asyncio +async def test_chat_json_with_data_prefix_is_not_streaming() -> None: + response = _json_response(_chat_json_with_data_prefix(), "application/json") + result, send_refund = await _settle_chat(response) + + send_refund.assert_awaited_once() + assert send_refund.await_args is not None + assert send_refund.await_args.args[0] == 6000 + assert result.headers["x-cashu"] == REFUND_TOKEN + assert result.headers["x-routstr-cost-msats"] == "4000" + body = json.loads(bytes(result.body)) + assert body["usage"]["cost"]["total_msats"] == 4000 + assert "data:image/png" in body["choices"][0]["message"]["content"] + + +@pytest.mark.asyncio +async def test_responses_json_with_data_prefix_is_not_streaming() -> None: + response = _json_response(_responses_json_with_data_prefix(), None) + result, send_refund = await _settle_responses(response) + + send_refund.assert_awaited_once() + assert send_refund.await_args is not None + assert send_refund.await_args.args[0] == 6000 + assert result.headers["x-cashu"] == REFUND_TOKEN + body = json.loads(bytes(result.body)) + assert body["usage"]["cost"]["total_msats"] == 4000 + + +@pytest.mark.asyncio +async def test_event_stream_content_type_is_streaming() -> None: + chunk = json.dumps( + { + "model": "gpt-5-mini", + "choices": [{"delta": {"content": "hi"}}], + "usage": {"prompt_tokens": 12, "completion_tokens": 8}, + } + ) + response = httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=f"data: {chunk}\n\ndata: [DONE]\n\n".encode(), + ) + + result, send_refund = await _settle_chat(response) + + send_refund.assert_awaited_once() + assert result.headers["x-cashu"] == REFUND_TOKEN + assert hasattr(result, "body_iterator") + + +@pytest.mark.parametrize( + ("content_type", "body", "expected"), + [ + ("text/event-stream", '{"a": 1}', True), + ("application/json", "data: {}\n\n", False), + ("application/json; charset=utf-8", 'data: "x"', False), + (None, '{"content": "data:image/png;base64,AAAA"}', False), + (None, "data: {}\n\n", True), + (None, "\n\n: keepalive\n\ndata: {}\n\n", True), + (None, "event: x\ndata: {}\n\n", True), + (None, '{"data:": 1}', False), + ("text/plain", "data: {}\n\n", True), + ("text/plain", '{"x": "data:"}', False), + (None, "", False), + ], +) +def test_is_sse_body(content_type: str | None, body: str, expected: bool) -> None: + assert _is_sse_body(content_type, body) is expected diff --git a/tests/unit/test_x_cashu_missing_usage.py b/tests/unit/test_x_cashu_missing_usage.py new file mode 100644 index 00000000..b2510d00 --- /dev/null +++ b/tests/unit/test_x_cashu_missing_usage.py @@ -0,0 +1,304 @@ +"""X-Cashu billing when the upstream omits usage. + +The local token estimator bills from the request body and the generated text. +When nothing can be estimated the prepayment is refunded in full. +""" + +import json +import logging +import os +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 + +REQUEST_BODY = json.dumps( + {"model": "gpt-4o", "messages": [{"role": "user", "content": "Tell me a joke"}]} +).encode() + + +def _sse(events: list[dict[str, Any]]) -> httpx.Response: + body = "".join(f"data: {json.dumps(e)}\n\n" for e in events) + "data: [DONE]\n\n" + return httpx.Response( + 200, headers={"content-type": "text/event-stream"}, content=body.encode() + ) + + +def _json(payload: dict[str, Any]) -> httpx.Response: + return httpx.Response( + 200, headers={"content-type": "application/json"}, content=json.dumps(payload) + ) + + +async def _settle( + response: httpx.Response, + *, + responses_api: bool = False, + request_body: bytes | None = REQUEST_BODY, + unit: str = "msat", +) -> tuple[Any, AsyncMock, AsyncMock]: + provider = BaseUpstreamProvider(base_url="http://test", api_key="test-key") + get_cost = AsyncMock(side_effect=provider.get_x_cashu_cost) + send_refund = AsyncMock(return_value="cashuBrefund") + handler = ( + provider.handle_x_cashu_responses_completion + if responses_api + else provider.handle_x_cashu_chat_completion + ) + with ( + patch.object(provider, "get_x_cashu_cost", new=get_cost), + patch.object(provider, "send_refund", new=send_refund), + ): + result = await handler( + response=response, + amount=10_000, + unit=unit, + max_cost_for_model=9_000, + mint=None, + request_body=request_body, + ) + return result, get_cost, send_refund + + +def _billed_usage(get_cost: AsyncMock) -> dict[str, Any] | None: + assert get_cost.await_args is not None + return get_cost.await_args.args[0].get("usage") + + +@pytest.mark.asyncio +async def test_streaming_chat_without_usage_bills_from_estimate() -> None: + events = [ + {"model": "gpt-4o", "choices": [{"delta": {"content": "Why did the "}}]}, + {"model": "gpt-4o", "choices": [{"delta": {"content": "chicken cross"}}]}, + ] + _, get_cost, _ = await _settle(_sse(events)) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage["input_tokens"] > 0 + assert usage["output_tokens"] > 0 + assert usage["estimated"] is True + + +@pytest.mark.asyncio +async def test_streaming_chat_without_text_refunds_everything() -> None: + _, get_cost, send_refund = await _settle( + _sse([{"model": "gpt-4o"}]), request_body=None + ) + + assert _billed_usage(get_cost) is None + assert send_refund.await_args is not None + assert send_refund.await_args.args[0] == 10_000 + + +@pytest.mark.asyncio +async def test_non_streaming_chat_without_usage_bills_from_estimate() -> None: + payload = { + "model": "gpt-4o", + "choices": [ + {"message": {"role": "assistant", "content": "To get to the other side."}} + ], + } + _, get_cost, _ = await _settle(_json(payload)) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage["output_tokens"] > 0 + assert usage["estimated"] is True + + +@pytest.mark.asyncio +async def test_streaming_responses_without_usage_bills_from_estimate() -> None: + events: list[dict[str, Any]] = [ + {"type": "response.created", "response": {"model": "gpt-5-mini"}}, + {"type": "response.output_text.delta", "delta": "Why did the chicken"}, + {"type": "response.output_text.done", "text": "Why did the chicken"}, + ] + _, get_cost, _ = await _settle(_sse(events), responses_api=True) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage["output_tokens"] > 0 + assert usage["estimated"] is True + + +@pytest.mark.asyncio +async def test_non_streaming_responses_without_usage_bills_from_estimate() -> None: + payload = { + "model": "gpt-5-mini", + "output": [ + {"type": "message", "content": [{"type": "output_text", "text": "Hi"}]} + ], + } + _, get_cost, _ = await _settle(_json(payload), responses_api=True) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage["output_tokens"] > 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("responses_api", [False, True]) +async def test_empty_streaming_usage_uses_estimate(responses_api: bool) -> None: + payload: dict[str, Any] = {"model": "gpt-4o", "usage": {}} + if responses_api: + payload["output"] = [{"content": [{"type": "output_text", "text": "Hello"}]}] + else: + payload["choices"] = [{"delta": {"content": "Hello"}}] + + _, get_cost, _ = await _settle(_sse([payload]), responses_api=responses_api) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage.get("estimated") is True + assert usage["output_tokens"] > 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("terminal_type", ["response.completed", "response.incomplete"]) +async def test_responses_terminal_output_is_not_billed_twice( + terminal_type: str, +) -> None: + payload = { + "model": "gpt-4o", + "output": [{"content": [{"type": "output_text", "text": "Hello world"}]}], + } + events: list[dict[str, Any]] = [ + {"type": "response.output_text.delta", "delta": "Hello world"}, + {"type": terminal_type, "response": payload}, + ] + _, streaming_cost, _ = await _settle(_sse(events), responses_api=True) + _, json_cost, _ = await _settle(_json(payload), responses_api=True) + + stream_usage = _billed_usage(streaming_cost) + json_usage = _billed_usage(json_cost) + assert stream_usage is not None and json_usage is not None + for field in ("input_tokens", "output_tokens", "total_tokens"): + assert stream_usage[field] == json_usage[field] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("input_shape", ["string", "message", "content_blocks"]) +async def test_responses_estimate_includes_input( + stream: bool, input_shape: str +) -> None: + prompt = "Explain how payment reservations work. " * 100 + response_input: Any = prompt + if input_shape == "message": + response_input = [{"role": "user", "content": prompt}] + elif input_shape == "content_blocks": + response_input = [ + {"role": "user", "content": [{"type": "input_text", "text": prompt}]} + ] + request_body = json.dumps( + { + "model": "gpt-4o", + "instructions": "Answer concisely.", + "input": response_input, + } + ).encode() + payload = { + "model": "gpt-4o", + "output": [{"content": [{"type": "output_text", "text": "Hello"}]}], + } + response = _sse([payload]) if stream else _json(payload) + _, get_cost, _ = await _settle( + response, responses_api=True, request_body=request_body + ) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage["input_tokens"] > 100 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("responses_api", [False, True]) +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("tokens", [0, 17]) +async def test_reported_usage_takes_precedence( + responses_api: bool, stream: bool, tokens: int +) -> None: + payload: dict[str, Any] = { + "model": "gpt-4o", + "usage": {"input_tokens": tokens, "output_tokens": tokens}, + } + if responses_api: + payload["output"] = [{"content": [{"type": "output_text", "text": "Hello"}]}] + else: + payload["choices"] = [{"message": {"content": "Hello"}}] + + response = _sse([payload]) if stream else _json(payload) + _, get_cost, _ = await _settle(response, responses_api=responses_api) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage["input_tokens"] == tokens + assert usage["output_tokens"] == tokens + assert "estimated" not in usage + + +@pytest.mark.asyncio +@pytest.mark.parametrize("responses_api", [False, True]) +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("unit", ["sat", "msat"]) +async def test_pricing_error_refunds_full_prepayment( + responses_api: bool, stream: bool, unit: str +) -> None: + payload = { + "model": "unpriced-model", + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + response = _sse([payload]) if stream else _json(payload) + with patch( + "routstr.payment.cost_calculation._get_pricing_rates", + side_effect=ValueError("No pricing for model"), + ): + result, _, send_refund = await _settle( + response, responses_api=responses_api, unit=unit + ) + + send_refund.assert_awaited_once_with(10_000, unit, None, request_id=None) + assert result.status_code == 200 + assert result.headers["X-Cashu"] == "cashuBrefund" + assert result.headers["X-Routstr-Cost-Msats"] == "0" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("responses_api", [False, True]) +async def test_full_refund_is_logged_with_model_provider_and_body( + responses_api: bool, caplog: pytest.LogCaptureFixture +) -> None: + payload = { + "model": "unpriced-model", + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + base_logger = logging.getLogger("routstr.upstream.base") + base_logger.addHandler(caplog.handler) + try: + with patch( + "routstr.payment.cost_calculation._get_pricing_rates", + side_effect=ValueError("No pricing for model"), + ): + await _settle(_json(payload), responses_api=responses_api) + finally: + base_logger.removeHandler(caplog.handler) + + record = next( + r + for r in caplog.records + if r.getMessage() == "Zero-cost settlement, refunding the full prepayment" + ) + logged = record.__dict__ + assert logged["model"] == "unpriced-model" + assert logged["provider_type"] == "base" + assert logged["upstream_base_url"] == "http://test" + assert logged["refund_amount"] == 10_000 + assert logged["unit"] == "msat" + assert "unpriced-model" in logged["response_body_preview"] diff --git a/tests/unit/test_x_cashu_provider_path.py b/tests/unit/test_x_cashu_provider_path.py new file mode 100644 index 00000000..11b47d6f --- /dev/null +++ b/tests/unit/test_x_cashu_provider_path.py @@ -0,0 +1,54 @@ +import json +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +from routstr.upstream.openrouter import OpenRouterUpstreamProvider + + +async def _body(response: Any) -> bytes: + chunks: list[bytes] = [] + async for chunk in response.body_iterator: + chunks.append(chunk if isinstance(chunk, bytes) else chunk.encode()) + return b"".join(chunks) + + +@pytest.mark.asyncio +async def test_x_cashu_chat_stream_reports_complete_provider_path() -> None: + provider = OpenRouterUpstreamProvider(api_key="test-key") + payload = {"model": "glm-4.5", "provider": "z.ai"} + content = f"data: {json.dumps(payload)}\n" + + 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, + ) + + event = json.loads((await _body(response)).decode().removeprefix("data: ")) + assert event["provider"] == "openrouter:z.ai" + + +@pytest.mark.asyncio +async def test_x_cashu_responses_stream_reports_complete_provider_path() -> None: + provider = OpenRouterUpstreamProvider(api_key="test-key") + event = {"type": "response.created", "provider": "z.ai"} + content = f"data: {json.dumps(event)}\n\n" + + 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, + ) + + payload = json.loads((await _body(response)).decode().removeprefix("data: ")) + assert payload["provider"] == "openrouter:z.ai" diff --git a/tests/unit/test_x_cashu_responses_streaming_sse.py b/tests/unit/test_x_cashu_responses_streaming_sse.py new file mode 100644 index 00000000..08154e40 --- /dev/null +++ b/tests/unit/test_x_cashu_responses_streaming_sse.py @@ -0,0 +1,221 @@ +"""X-Cashu settlement for streaming ``/v1/responses``. + +The stream is real SSE: CRLF delimiters, comment keepalives, ``event:`` fields, +multi-line ``data:`` payloads and a ``[DONE]`` sentinel. Canonical Responses API +usage arrives nested under ``response`` on ``response.completed``. +""" + +import json +import os +from typing import Any +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.payment.cost_calculation import CostData # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 + + +def _make_provider() -> BaseUpstreamProvider: + return BaseUpstreamProvider(base_url="http://test", api_key="test-key") + + +def _make_cost_data(total_msats: int = 4000) -> CostData: + return CostData( + base_msats=0, + input_msats=2500, + output_msats=1500, + total_msats=total_msats, + total_usd=0.0002, + input_tokens=12, + output_tokens=8, + ) + + +def _sse_response(chunks: list[bytes]) -> httpx.Response: + """Build the upstream response from wire chunks that split events.""" + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=b"".join(chunks), + ) + + +COMPLETED_EVENT = { + "type": "response.completed", + "response": { + "model": "gpt-5-mini", + "usage": { + "input_tokens": 12, + "output_tokens": 8, + "total_tokens": 20, + "output_tokens_details": {"reasoning_tokens": 3}, + }, + }, +} + + +def _canonical_chunks() -> list[bytes]: + """CRLF stream whose completed event straddles two wire chunks.""" + completed = json.dumps(COMPLETED_EVENT).encode() + return [ + b": keepalive\r\n\r\n", + b"event: response.created\r\n" + b'data: {"type":"response.created","response":{"model":"gpt-5-mini"}}\r\n\r\n', + b"event: response.completed\r\ndata: " + completed[:40], + completed[40:] + b"\r\n\r\n", + b"data: [DONE]\r\n\r\n", + ] + + +async def _collect(response: Any) -> bytes: + body = b"" + async for chunk in response.body_iterator: + body += chunk + return body + + +async def _settle( + chunks: list[bytes], + *, + amount: int = 10_000, + max_cost_for_model: int = 9_000, + cost_data: CostData | None = None, +) -> tuple[Any, AsyncMock, AsyncMock]: + provider = _make_provider() + get_cost = ( + AsyncMock(return_value=cost_data) + if cost_data is not None + else AsyncMock(side_effect=provider.get_x_cashu_cost) + ) + send_refund = AsyncMock(return_value="cashuBrefundtoken0123456789") + with ( + patch.object(provider, "get_x_cashu_cost", new=get_cost), + patch.object(provider, "send_refund", new=send_refund), + ): + response = await provider.handle_x_cashu_responses_completion( + response=_sse_response(chunks), + amount=amount, + unit="msat", + max_cost_for_model=max_cost_for_model, + mint=None, + ) + return response, get_cost, send_refund + + +@pytest.mark.asyncio +async def test_fragmented_crlf_stream_refunds_and_sets_cost_headers() -> None: + response, _, send_refund = await _settle( + _canonical_chunks(), cost_data=_make_cost_data(4000) + ) + + send_refund.assert_awaited_once() + assert send_refund.await_args is not None + assert send_refund.await_args.args[0] == 10_000 - 4000 + assert response.headers["x-cashu"] == "cashuBrefundtoken0123456789" + assert response.headers["x-routstr-cost-msats"] == "4000" + assert response.headers["x-routstr-input-cost-msats"] == "2500" + assert response.headers["x-routstr-output-cost-msats"] == "1500" + + +@pytest.mark.asyncio +async def test_nested_completion_usage_drives_cost_calculation() -> None: + _, get_cost, _ = await _settle(_canonical_chunks(), cost_data=_make_cost_data(4000)) + + assert get_cost.await_args is not None + response_data = get_cost.await_args.args[0] + assert response_data["model"] == "gpt-5-mini" + assert response_data["usage"]["input_tokens"] == 12 + assert response_data["usage"]["output_tokens"] == 8 + + +@pytest.mark.asyncio +async def test_reemitted_stream_is_valid_sse() -> None: + response, _, _ = await _settle(_canonical_chunks(), cost_data=_make_cost_data(4000)) + body = await _collect(response) + + assert b"\\n" not in body + assert body.endswith(b"\n\n") + assert b": keepalive" not in body + + events = [e for e in body.split(b"\n\n") if e.strip()] + payloads = [] + for event in events: + data_lines = [ + line[len(b"data:") :].lstrip() + for line in event.split(b"\n") + if line.startswith(b"data:") + ] + assert data_lines, f"event carries no data line: {event!r}" + payloads.append(b"\n".join(data_lines)) + + assert payloads[-1] == b"[DONE]" + assert any(b"event: response.completed" in event for event in events) + + completed = json.loads(payloads[-2]) + assert completed["type"] == "response.completed" + assert completed["response"]["usage"]["cost"]["total_msats"] == 4000 + + +@pytest.mark.asyncio +async def test_multiline_data_payload_is_parsed_and_reframed() -> None: + completed = json.dumps(COMPLETED_EVENT) + head, tail = completed[:30], completed[30:] + chunks = [ + ("data: " + head + "\r\ndata: " + tail + "\r\n\r\n").encode(), + b"data: [DONE]\r\n\r\n", + ] + + response, get_cost, send_refund = await _settle( + chunks, cost_data=_make_cost_data(4000) + ) + + assert get_cost.await_args is not None + assert send_refund.await_args is not None + assert get_cost.await_args.args[0]["usage"]["input_tokens"] == 12 + assert send_refund.await_args.args[0] == 6000 + body = await _collect(response) + for event in body.split(b"\n\n"): + for line in event.split(b"\n"): + if line.strip(): + assert line.startswith(b"data:") or line.startswith(b"event:") + + +@pytest.mark.asyncio +async def test_missing_usage_refunds_instead_of_charging_authorized_max() -> None: + chunks = [ + b'data: {"type":"response.created","response":{"model":"gpt-5-mini"}}\r\n\r\n', + b"data: [DONE]\r\n\r\n", + ] + + response, _, send_refund = await _settle( + chunks, amount=10_000, max_cost_for_model=9_000 + ) + + send_refund.assert_awaited_once() + assert send_refund.await_args is not None + assert send_refund.await_args.args[0] == 10_000 + assert response.headers["x-cashu"] == "cashuBrefundtoken0123456789" + assert response.headers["x-routstr-cost-msats"] == "0" + + +@pytest.mark.asyncio +async def test_malformed_events_do_not_retain_whole_token() -> None: + chunks = [ + b"data: {not json\r\n\r\n", + b"data: [DONE]\r\n\r\n", + ] + + response, _, send_refund = await _settle( + chunks, amount=10_000, max_cost_for_model=9_000 + ) + + assert send_refund.await_args is not None + assert send_refund.await_args.args[0] == 10_000 + body = await _collect(response) + assert b"\\n" not in body + assert body.endswith(b"\n\n") diff --git a/ui/app/transactions/page.tsx b/ui/app/transactions/page.tsx index 0ea721c0..56bc1b74 100644 --- a/ui/app/transactions/page.tsx +++ b/ui/app/transactions/page.tsx @@ -3,6 +3,7 @@ import { useState, useEffect } from 'react'; import { useQuery, keepPreviousData } from '@tanstack/react-query'; import { useCopyToClipboard } from '@/hooks/use-copy-to-clipboard'; +import { downloadText } from '@/lib/utils'; import { AppPageShell } from '@/components/app-page-shell'; import { PageHeader } from '@/components/page-header'; import { @@ -53,6 +54,7 @@ import { Zap, ChevronLeft, ChevronRight, + Download, } from 'lucide-react'; import { AdminService, @@ -64,6 +66,28 @@ import { toast } from 'sonner'; const STORAGE_KEY = 'routstr-transaction-filters'; +const STATUS_BADGES: Record< + Transaction['status'], + { label: string; className: string } +> = { + issued: { + label: 'Issued', + className: 'border-gray-500/20 bg-gray-500/10 text-gray-500', + }, + collected: { + label: 'Collected', + className: 'border-green-500/20 bg-green-500/10 text-green-500', + }, + swept: { + label: 'Swept', + className: 'border-orange-500/20 bg-orange-500/10 text-orange-500', + }, + pending: { + label: 'Pending', + className: 'border-blue-500/20 bg-blue-500/10 text-blue-500', + }, +}; + function TransactionTable({ transactions, copiedId, @@ -181,19 +205,36 @@ function TransactionTable({ {format(tx.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')} - + {/* The token is never rendered in the table, so the + clipboard must not be the only way to get it out. */} + {tx.status === 'issued' && ( + )} - + ))} @@ -402,6 +443,7 @@ export default function TransactionsPage() { const [activeTab, setActiveTab] = useState('x-cashu'); const [xcashuPage, setXcashuPage] = useState(0); const [apikeyPage, setApikeyPage] = useState(0); + const [withdrawalsPage, setWithdrawalsPage] = useState(0); const [lightningPage, setLightningPage] = useState(0); const typeParam = type === 'all' ? undefined : type; @@ -450,6 +492,28 @@ export default function TransactionsPage() { placeholderData: keepPreviousData, }); + // Withdrawals are stored with source "admin" and keep their one-time token. + const withdrawalsQuery = useQuery({ + queryKey: [ + 'transactions', + 'admin', + typeParam, + statusParam, + searchParam, + withdrawalsPage, + ], + queryFn: () => + AdminService.getTransactions( + typeParam, + statusParam, + searchParam, + 'admin', + PAGE_SIZE, + withdrawalsPage * PAGE_SIZE + ), + placeholderData: keepPreviousData, + }); + const LIGHTNING_STATUSES = ['pending', 'paid', 'expired', 'cancelled']; const lightningStatusParam = LIGHTNING_STATUSES.includes(status) ? status @@ -480,6 +544,7 @@ export default function TransactionsPage() { setStatus('all'); setXcashuPage(0); setApikeyPage(0); + setWithdrawalsPage(0); setLightningPage(0); }; @@ -490,30 +555,10 @@ export default function TransactionsPage() { }; const getStatusBadge = (tx: Transaction) => { - if (tx.swept) - return ( - - Swept - - ); - if (tx.collected) - return ( - - Collected - - ); + const { label, className } = STATUS_BADGES[tx.status]; return ( - - Pending + + {label} ); }; @@ -533,12 +578,14 @@ export default function TransactionsPage() { useEffect(() => { setXcashuPage(0); setApikeyPage(0); + setWithdrawalsPage(0); setLightningPage(0); }, [type, status, search]); const isRefetching = xcashuQuery.isRefetching || apikeyQuery.isRefetching || + withdrawalsQuery.isRefetching || lightningQuery.isRefetching; const renderCardContent = ( @@ -617,6 +664,7 @@ export default function TransactionsPage() { onClick={() => { xcashuQuery.refetch(); apikeyQuery.refetch(); + withdrawalsQuery.refetch(); lightningQuery.refetch(); }} variant='outline' @@ -675,6 +723,7 @@ export default function TransactionsPage() { All Statuses Pending + Issued Collected Swept Paid (Lightning) @@ -703,7 +752,7 @@ export default function TransactionsPage() { value={activeTab} onValueChange={setActiveTab} > - + X-Cashu @@ -722,6 +771,18 @@ export default function TransactionsPage() { )} + + + Withdrawals + {withdrawalsQuery.data && ( + + {withdrawalsQuery.data.total} + + )} + Lightning @@ -769,6 +830,28 @@ export default function TransactionsPage() { + + + +
+ Withdrawal History + {hasActiveFilters && ( + + Filtered by {activeFilterDescription} + + )} +
+
+ + {renderCardContent( + withdrawalsQuery, + withdrawalsPage, + setWithdrawalsPage + )} + +
+
+ diff --git a/ui/components/child-key-creator.tsx b/ui/components/child-key-creator.tsx deleted file mode 100644 index b10847f9..00000000 --- a/ui/components/child-key-creator.tsx +++ /dev/null @@ -1,481 +0,0 @@ -'use client'; - -import { useState } from 'react'; -import { useCopyToClipboard } from '@/hooks/use-copy-to-clipboard'; -import { useWalletInfo } from '@/hooks/use-wallet-info'; -import { WalletService } from '@/lib/api/services/wallet'; -import { ApiKeyInput } from './api-key-input'; -import { Button } from '@/components/ui/button'; -import { - Card, - CardContent, - CardDescription, - CardHeader, - CardTitle, -} from '@/components/ui/card'; -import { Alert, AlertDescription, AlertTitle } from '@/components/ui/alert'; -import { Input } from '@/components/ui/input'; -import { Textarea } from '@/components/ui/textarea'; -import { Label } from '@/components/ui/label'; -import { Key, Copy, Check, Loader2, Plus, Trash2 } from 'lucide-react'; -import { toast } from 'sonner'; -import { KeyOptions } from './key-options'; - -interface KeyConfig { - id: string; - count: number; - balanceLimit: string; - balanceLimitReset: string; - validityDate: string; -} - -interface ChildKeyCreatorProps { - baseUrl?: string; - apiKey?: string; - onApiKeyChange?: (apiKey: string) => void; - costPerKeyMsats?: number; -} - -function formatSats(msats: number): string { - return new Intl.NumberFormat('en-US').format(Math.floor(msats / 1000)); -} - -function formatMsats(msats: number): string { - return new Intl.NumberFormat('en-US').format(msats); -} - -export function ChildKeyCreator({ - baseUrl, - apiKey: propApiKey, - onApiKeyChange, - costPerKeyMsats, -}: ChildKeyCreatorProps) { - const [internalApiKey, setInternalApiKey] = useState(''); - const [loading, setLoading] = useState(false); - const [error, setError] = useState(null); - const [configs, setConfigs] = useState([ - { - id: crypto.randomUUID(), - count: 1, - balanceLimit: '', - balanceLimitReset: '', - validityDate: '', - }, - ]); - - const activeApiKey = propApiKey ?? internalApiKey; - const { data: walletInfo } = useWalletInfo(baseUrl ?? '', activeApiKey); - - const handleApiKeyChange = (val: string) => { - setInternalApiKey(val); - onApiKeyChange?.(val); - }; - - const [newKeys, setNewKeys] = useState([]); - const [resultInfo, setResultInfo] = useState<{ - cost_msats: number; - parent_balance: number; - } | null>(null); - const { copiedKey, copy } = useCopyToClipboard(); - - const addConfig = () => { - setConfigs([ - ...configs, - { - id: crypto.randomUUID(), - count: 1, - balanceLimit: '', - balanceLimitReset: '', - validityDate: '', - }, - ]); - }; - - const removeConfig = (id: string) => { - if (configs.length > 1) { - setConfigs(configs.filter((c) => c.id !== id)); - } - }; - - const updateConfig = (id: string, updates: Partial) => { - setConfigs(configs.map((c) => (c.id === id ? { ...c, ...updates } : c))); - }; - - const handleCreateKey = async () => { - if (!activeApiKey && baseUrl) { - toast.error('Please provide a Parent API key first'); - return; - } - - setLoading(true); - setError(null); - try { - let allNewKeys: string[] = []; - let totalCost = 0; - let lastParentBalance = 0; - - for (const config of configs) { - const requestedCount = Math.max(1, Math.min(50, Number(config.count))); - const result = await WalletService.createChildKey( - baseUrl, - activeApiKey, - requestedCount, - config.balanceLimit ? parseInt(config.balanceLimit) : undefined, - config.balanceLimitReset || undefined, - config.validityDate - ? Math.floor( - new Date(config.validityDate + 'T23:59:59').getTime() / 1000 - ) - : undefined - ); - - if (result.api_keys) { - allNewKeys = [...allNewKeys, ...result.api_keys]; - } - totalCost += result.cost_msats; - lastParentBalance = result.parent_balance; - } - - setNewKeys(allNewKeys); - setResultInfo({ - cost_msats: totalCost, - parent_balance: lastParentBalance, - }); - - toast.success( - `${allNewKeys.length} child API key${ - allNewKeys.length > 1 ? 's' : '' - } created successfully` - ); - } catch (error) { - console.error('Failed to create child key:', error); - let errorMessage = - error instanceof Error ? error.message : 'Failed to create child key'; - try { - const parsed = JSON.parse(errorMessage); - errorMessage = - parsed.detail?.error?.message || - (typeof parsed.detail === 'string' ? parsed.detail : errorMessage); - } catch {} - setError(errorMessage); - toast.error(errorMessage); - } finally { - setLoading(false); - } - }; - - const copyToClipboard = async (key: string) => { - if (await copy(key, key)) { - toast.success('API key copied to clipboard'); - } - }; - - const copyAllToClipboard = async () => { - if (await copy(newKeys.join('\n'), 'all')) { - toast.success('All API keys copied to clipboard'); - } - }; - - return ( -
- - -
-
- Create Child API Key - - Generate secondary API keys that share your account balance. - -
- {costPerKeyMsats !== undefined && ( -
-

- Unit Cost -

-

- {costPerKeyMsats / 1000} sats -

-
- )} -
-
- -
- {baseUrl && ( -
- -
-
- -
- -
- {walletInfo && ( -
-
- - Spendable Balance - - - {formatSats(walletInfo.balanceMsats)} sats - -
-
- - Total Requests - - - {walletInfo.totalRequests} - -
-
- - Total Spent - -
-

- {formatSats(walletInfo.totalSpent)} sats -

-

- {formatMsats(walletInfo.totalSpent)} msats -

-
-
-
- )} -
- )} - -
- {configs.map((config) => ( -
- {configs.length > 1 && ( - - )} -
-
- - { - const val = parseInt(e.target.value); - updateConfig(config.id, { - count: isNaN(val) - ? 1 - : Math.max(1, Math.min(50, val)), - }); - }} - className='h-9' - /> -
- -
- - updateConfig(config.id, { balanceLimit: val }) - } - validityDate={config.validityDate} - setValidityDate={(val) => - updateConfig(config.id, { validityDate: val }) - } - balanceLimitReset={config.balanceLimitReset} - setBalanceLimitReset={(val) => - updateConfig(config.id, { balanceLimitReset: val }) - } - /> -
-
-
- ))} - -
- -
- -
-
- {costPerKeyMsats && ( -

- Total Cost:{' '} - - {costPerKeyMsats * - configs.reduce( - (acc, c) => acc + Number(c.count), - 0 - )}{' '} - mSats - -

- )} -
- - -
- - {error && ( - - Error - {error} - - )} - -

- Each key creation has a small one-time fee. -

-
- - {newKeys.length > 0 && ( -
- - - {newKeys.length} New API Key{newKeys.length > 1 ? 's' : ''}{' '} - Generated - - - Copy {newKeys.length > 1 ? 'these keys' : 'this key'} now. - You won't be able to see them again. - {resultInfo && ( -
- Total Cost: {resultInfo.cost_msats / 1000} sats | New - Balance: {resultInfo.parent_balance / 1000} sats -
- )} -
-
- -
-
- - Generated Keys ({newKeys.length}) - - {newKeys.length > 1 && ( - - )} -
-
- {newKeys.map((key, index) => ( -
- - {key} - - -
- ))} -
-
- - {newKeys.length > 3 && ( -
- -
-