mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-30 23:36:15 +00:00
Compare commits
135
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7aa3db4c68 | ||
|
|
2efbc3783c | ||
|
|
2cf2592cce | ||
|
|
094e7195f0 | ||
|
|
a9593ab416 | ||
|
|
6221ee8152 | ||
|
|
65ea28cb85 | ||
|
|
88fe9758a3 | ||
|
|
03b00f3eb6 | ||
|
|
7e8c2033ae | ||
|
|
52742a6a04 | ||
|
|
030d8b61ce | ||
|
|
f47e16aa61 | ||
|
|
eb6a612e5a | ||
|
|
c2a2d76eae | ||
|
|
d4657dca5a | ||
|
|
7106dfe330 | ||
|
|
1045f1061b | ||
|
|
0415806c5a | ||
|
|
5d5c849180 | ||
|
|
b8700dde40 | ||
|
|
4872c318d5 | ||
|
|
56a67c0a86 | ||
|
|
afcb3f7cda | ||
|
|
14748c28fb | ||
|
|
8aa2fa5c4a | ||
|
|
a024d5be5e | ||
|
|
c40723c1e2 | ||
|
|
6c5c3f0b5f | ||
|
|
86adee9b1b | ||
|
|
18965b4ea4 | ||
|
|
fcc87718ff | ||
|
|
50437a1cc6 | ||
|
|
7f5a0cf1ae | ||
|
|
3f850c54f4 | ||
|
|
df4d4c44e6 | ||
|
|
02cd2cfeec | ||
|
|
0aebfc6dbe | ||
|
|
b76fa17f81 | ||
|
|
ed8ac914f9 | ||
|
|
6c762f0c3d | ||
|
|
2b4f70442a | ||
|
|
99724cc5f5 | ||
|
|
e19f679609 | ||
|
|
7481ebd8cd | ||
|
|
7be1db8fd9 | ||
|
|
4a4443f98f | ||
|
|
fecf5b27b8 | ||
|
|
bd1edcef26 | ||
|
|
22a94a68a4 | ||
|
|
a3a4d69ed3 | ||
|
|
c9533c872a | ||
|
|
90da3803c6 | ||
|
|
8b3b59e176 | ||
|
|
be5e323c68 | ||
|
|
f6d1a41728 | ||
|
|
b3bf1f0e90 | ||
|
|
3035292d2a | ||
|
|
ac436d0862 | ||
|
|
59bcc3cbbf | ||
|
|
25f75643c2 | ||
|
|
c68c1936d3 | ||
|
|
49d285571e | ||
|
|
eb108a4a5a | ||
|
|
19082231f9 | ||
|
|
281607108c | ||
|
|
8314f3c1b0 | ||
|
|
1a9041766b | ||
|
|
01c01fe8ad | ||
|
|
892aed61cc | ||
|
|
c6733dbb62 | ||
|
|
57fef44ce7 | ||
|
|
c671d277d3 | ||
|
|
9460e24f21 | ||
|
|
f31929b538 | ||
|
|
129ea7bb76 | ||
|
|
b88c6d84fa | ||
|
|
238beb3e52 | ||
|
|
dc833d9f5f | ||
|
|
31e7ff8fbd | ||
|
|
a1850dfa55 | ||
|
|
9eacabde0e | ||
|
|
73f52ef1fd | ||
|
|
110ec5da5d | ||
|
|
a5cee796c9 | ||
|
|
aabf08a930 | ||
|
|
07c15de106 | ||
|
|
aea24d23da | ||
|
|
e100c77288 | ||
|
|
057a752b1e | ||
|
|
5daa2602f5 | ||
|
|
d80912b10e | ||
|
|
8f087bf07a | ||
|
|
744321153f | ||
|
|
b81c5add6a | ||
|
|
02109616b9 | ||
|
|
b30346bcb6 | ||
|
|
956a1ac3e1 | ||
|
|
4287f038cf | ||
|
|
82fd2c08a7 | ||
|
|
a1118a510d | ||
|
|
27dedfdf77 | ||
|
|
7c2e2ce512 | ||
|
|
89b93fde98 | ||
|
|
0f7d3d2f86 | ||
|
|
af2abc3a1c | ||
|
|
9b4b4d734f | ||
|
|
edfb0f127c | ||
|
|
bc6d0af6d6 | ||
|
|
6097193442 | ||
|
|
65ac1b0dcd | ||
|
|
ab5596ed70 | ||
|
|
55622cb956 | ||
|
|
7328a5ac45 | ||
|
|
07d39c2a7b | ||
|
|
bd4f4e8207 | ||
|
|
46b1f8f72d | ||
|
|
1c0570cc0e | ||
|
|
a8f60643e0 | ||
|
|
e4753b5864 | ||
|
|
66d7bd87b5 | ||
|
|
4c48df8aa8 | ||
|
|
abc1ea5b35 | ||
|
|
52dd011cd8 | ||
|
|
94a3215894 | ||
|
|
5dd3cbbc22 | ||
|
|
55fc25b4de | ||
|
|
36ed3ec9f0 | ||
|
|
2e350e082b | ||
|
|
24b90af6e6 | ||
|
|
b9bf0ea23b | ||
|
|
11fc826bfb | ||
|
|
757d4400af | ||
|
|
766d59d666 | ||
|
|
98ddd37853 |
+17
-2
@@ -2,7 +2,23 @@
|
||||
UPSTREAM_BASE_URL=https://api.openai.com/v1
|
||||
UPSTREAM_API_KEY=your-upstream-api-key
|
||||
|
||||
# ADMIN_PASSWORD=secure-admin-password
|
||||
# Tinfoil (confidential inference enclaves, EHBP)
|
||||
# TINFOIL_API_KEY=your-tinfoil-api-key
|
||||
|
||||
# Secret key used to encrypt node secrets at rest (optional). If unset, the node
|
||||
# generates one on first start, writes it to routstr_secret.key (override the path
|
||||
# with ROUTSTR_SECRET_KEY_FILE), and prints it once — back that file up, because
|
||||
# losing the key makes previously encrypted secrets unreadable. Set it explicitly
|
||||
# to manage the key yourself (recommended in production). Generate one with:
|
||||
# uv run python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())"
|
||||
ROUTSTR_SECRET_KEY=
|
||||
|
||||
# The admin password and the Nostr identity (nsec) are NOT set here. The admin
|
||||
# password is generated and logged once on first start (read it from the logs to
|
||||
# sign in); both are managed afterwards from the admin UI and stored encrypted in
|
||||
# the database. ADMIN_PASSWORD / NSEC are still read once as a legacy seed for
|
||||
# existing deployments, but new nodes should set them in the UI — a value left in
|
||||
# .env is ignored once the node has been configured.
|
||||
|
||||
# Database
|
||||
# DATABASE_URL=sqlite+aiosqlite:///keys.db
|
||||
@@ -10,7 +26,6 @@ UPSTREAM_API_KEY=your-upstream-api-key
|
||||
# Node Information
|
||||
# NAME=My Routstr Node
|
||||
# DESCRIPTION=Fast AI API access with Bitcoin payments
|
||||
# NSEC=nsec1...
|
||||
# HTTP_URL=https://api.mynode.com
|
||||
# ONION_URL=http://mynode.onion (auto fetched from compose)
|
||||
# RELAYS="wss://relay.damus.io,wss://relay.nostr.band,wss://eden.nostr.land,wss://relay.routstr.com"
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
__pycache__
|
||||
.env
|
||||
keys.db
|
||||
routstr_secret.key
|
||||
wallet.sqlite3
|
||||
|
||||
# Python build artifacts
|
||||
@@ -16,6 +17,9 @@ dist/
|
||||
*.db-shm
|
||||
*.db-wal
|
||||
.*wallet.sqlite3
|
||||
.wallet/
|
||||
AGENTS.md
|
||||
TEST_SUITE_OVERVIEW.md
|
||||
*models.json
|
||||
.cashu
|
||||
.relay
|
||||
@@ -38,3 +42,4 @@ proof_backups
|
||||
|
||||
*.todo
|
||||
ui_out
|
||||
.worktrees
|
||||
|
||||
+1
-1
@@ -7,7 +7,7 @@ WORKDIR /app/ui
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
|
||||
|
||||
# Copy UI source
|
||||
COPY ui/package.json ui/pnpm-lock.yaml* ./
|
||||
COPY ui/package.json ui/pnpm-lock.yaml* ui/pnpm-workspace.yaml* ./
|
||||
RUN pnpm install --frozen-lockfile
|
||||
|
||||
COPY ui/ ./
|
||||
|
||||
@@ -55,19 +55,41 @@ If you are a node runner, start a Routstr Core instance using Docker Compose:
|
||||
|
||||
1. **Prepare your `.env`**:
|
||||
```bash
|
||||
ADMIN_PASSWORD=mysecretpassword
|
||||
# Optional: encrypts node secrets at rest. If unset, the node generates a key
|
||||
# on first start, writes it to routstr_secret.key, and prints it once — back
|
||||
# up that file. Set it explicitly to manage the key yourself (recommended in
|
||||
# production).
|
||||
ROUTSTR_SECRET_KEY=<generated-key>
|
||||
NAME="My AI Node"
|
||||
DESCRIPTION="Fast access to models"
|
||||
NSEC=yournsec
|
||||
RECEIVE_LN_ADDRESS=yourname@wallet.com
|
||||
```
|
||||
|
||||
Your Nostr identity (`nsec`) is not set in `.env` — configure it from the admin
|
||||
UI after first start, where it's stored encrypted in the database. (`NSEC` in
|
||||
`.env` is still read once as a legacy seed for existing deployments.)
|
||||
|
||||
If you don't set one, a key is generated and printed on first start — save it
|
||||
somewhere safe (losing it makes previously encrypted secrets unreadable). To
|
||||
supply your own, generate it once and keep it stable:
|
||||
```bash
|
||||
uv run python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())"
|
||||
```
|
||||
|
||||
2. **Start the services**:
|
||||
```bash
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
3. **Configure**:
|
||||
3. **Get your admin password**:
|
||||
On first start the node generates an admin password and logs it once with the
|
||||
`/admin` URL. Read it from the logs:
|
||||
```bash
|
||||
docker compose logs routstr | grep -i admin
|
||||
```
|
||||
(Lost it? Reset with `docker compose exec routstr /.venv/bin/python scripts/reset_admin_password.py --regenerate`.)
|
||||
|
||||
4. **Configure**:
|
||||
Open [http://localhost:8000/admin/](http://localhost:8000/admin/) to connect your AI providers and set pricing.
|
||||
|
||||
For full instructions, see the **[Provider Quick Start Guide](https://docs.routstr.com/provider/quickstart/)**.
|
||||
|
||||
@@ -13,6 +13,7 @@ services:
|
||||
- ./ui_out:/output:z
|
||||
command:
|
||||
["sh", "-c", "mkdir -p /output && cp -r /app/built/. /output/ && echo 'UI build copied to mounted volume' && ls -la /output/ && echo 'UI built and ready' && tail -f /dev/null"]
|
||||
restart: unless-stopped
|
||||
|
||||
routstr:
|
||||
build: .
|
||||
@@ -31,6 +32,7 @@ services:
|
||||
- 8000:8000
|
||||
extra_hosts: # Needed to access locally running models
|
||||
- "host.docker.internal:host-gateway"
|
||||
restart: unless-stopped
|
||||
|
||||
tor:
|
||||
image: ghcr.io/hundehausen/tor-hidden-service:latest
|
||||
@@ -41,6 +43,7 @@ services:
|
||||
- HS_ROUTER=routstr:8000:80
|
||||
depends_on:
|
||||
- routstr
|
||||
restart: unless-stopped
|
||||
|
||||
volumes:
|
||||
tor-data:
|
||||
|
||||
@@ -0,0 +1,150 @@
|
||||
# EHBP Proxy Support for Tinfoil Models
|
||||
|
||||
## Problem
|
||||
|
||||
The SDK's `SecureClient.fetch` encrypts request bodies with HPKE (EHBP protocol)
|
||||
and sends them to the Routstr provider. The Routstr proxy had no EHBP handling:
|
||||
|
||||
1. It tried to `json.loads()` the binary HPKE-sealed body → failed with a 400
|
||||
2. The upstream PPQ.AI public endpoint (`/v1/chat/completions`) doesn't speak
|
||||
EHBP, so the response had no `Ehbp-Response-Nonce` header
|
||||
3. `SecureClient` threw `Missing Ehbp-Response-Nonce header` because it expects
|
||||
every response from an EHBP-configured `baseURL` to carry that header
|
||||
|
||||
## Root cause
|
||||
|
||||
PPQ.AI exposes EHBP-aware inference at `/private/`, separate from the public
|
||||
`/v1/` endpoint. The Routstr proxy was forwarding to `/v1/` (the public
|
||||
endpoint) instead of `/private/` (the enclave endpoint). The public endpoint
|
||||
can't decrypt the body, returns a normal HTTP response, and the SDK can't
|
||||
decrypt it because there's no nonce header.
|
||||
|
||||
The PPQ private-mode proxy (`ppq-private-mode-proxy/lib/proxy.ts`) shows the
|
||||
correct pattern: `SecureClient` talks to `api.ppq.ai/private/v1/chat/completions`,
|
||||
which decrypts inside the attested enclave and returns an EHBP-encrypted
|
||||
response with the `Ehbp-Response-Nonce` header.
|
||||
|
||||
## What was changed
|
||||
|
||||
### `routstr/proxy.py`
|
||||
|
||||
Detects EHBP requests by checking for the `Ehbp-Encapsulated-Key` header (set
|
||||
by the EHBP transport on every encrypted request). For EHBP requests:
|
||||
|
||||
- Skips JSON body parsing (the body is binary ciphertext, not JSON)
|
||||
- Reads the model ID from the `X-Routstr-Model` header (set by the SDK) instead
|
||||
of from `body.model`
|
||||
- Routes through new `forward_ehbp_request` (bearer auth) and
|
||||
`forward_ehbp_x_cashu_request` (x-cashu auth) methods
|
||||
- Skips reactive 400 param correction (can't parse encrypted response body)
|
||||
- Still charges the user via `pay_for_request` (uses `max_cost_for_model` from
|
||||
the model registry, not the body)
|
||||
|
||||
### `routstr/upstream/base.py`
|
||||
|
||||
Keeps EHBP as an explicit opt-in provider capability instead of making every
|
||||
upstream provider appear EHBP-capable:
|
||||
|
||||
- `supports_ehbp = False` by default
|
||||
- `get_ehbp_forwarding_target(path, model_obj)` raises `NotImplementedError`
|
||||
unless a provider opts in and returns a provider-specific EHBP target
|
||||
|
||||
The actual EHBP forwarding logic does **not** live in `base.py`.
|
||||
|
||||
### `routstr/upstream/ehbp.py`
|
||||
|
||||
Contains the shared opaque EHBP transport and billing helpers:
|
||||
|
||||
- `EHBPForwardingTarget` — provider-specific target URL plus extra headers
|
||||
- `forward_ehbp_request()` — forwards the encrypted body, captures Tinfoil
|
||||
usage from a response header or streaming HTTP trailer, and finalizes bearer
|
||||
billing at actual cost (falling back to max cost when usage is unavailable)
|
||||
- `forward_ehbp_x_cashu_request()` — redeems the Cashu token, refunds the full
|
||||
token on upstream failure, and refunds the difference between the redeemed
|
||||
amount and actual cost (or max cost when usage is unavailable)
|
||||
|
||||
### Provider support
|
||||
|
||||
EHBP is currently enabled only for `TinfoilUpstreamProvider`. It forwards to
|
||||
Tinfoil's attested enclave and requests `X-Tinfoil-Usage-Metrics` for billing.
|
||||
PPQ.AI retains its private-target implementation, but `supports_ehbp = False`
|
||||
until it has a provider-specific trusted usage/model-binding strategy.
|
||||
|
||||
## Why it's done this way
|
||||
|
||||
The proxy is a **blind relay** for EHBP requests. It cannot decrypt the body
|
||||
(only the attested enclave can), so it must:
|
||||
|
||||
1. Get the model ID from a header, not the body
|
||||
2. Forward the raw bytes without parsing or transformation
|
||||
3. Stream the response back without SSE/cost parsing
|
||||
4. Pass through EHBP protocol headers (`Ehbp-Encapsulated-Key` on request,
|
||||
`Ehbp-Response-Nonce` on response)
|
||||
|
||||
Cost tracking happens at the proxy level. Routstr reserves or redeems up to
|
||||
`max_cost_for_model`, then Tinfoil's out-of-band usage header/trailer allows it
|
||||
to finalize at actual token cost. If trusted usage is missing or invalid, the
|
||||
proxy safely falls back to max-cost billing.
|
||||
|
||||
## End-to-end flow
|
||||
|
||||
```
|
||||
SDK Routstr Proxy PPQ.AI /private/
|
||||
│ │ │
|
||||
│── X-Routstr-Model: tinfoil-kimi-k2-6 ─│ │
|
||||
│── Ehbp-Encapsulated-Key: <hex> ───────│ │
|
||||
│── Authorization: Bearer <cashu> ──────│ │
|
||||
│── body = HPKE-encrypted(kimi-k2-6) ───│ │
|
||||
│ │ │
|
||||
│ detects Ehbp-Encapsulated-Key │
|
||||
│ reads model from X-Routstr-Model │
|
||||
│ does billing/routing │
|
||||
│ │ │
|
||||
│ adds X-Private-Model: private/kimi-k2-6
|
||||
│ forwards raw body to /private/v1/... │
|
||||
│ │──────────────────────────────▶│
|
||||
│ │ enclave decrypts
|
||||
│ │ runs inference
|
||||
│ │◀── Ehbp-Response-Nonce ──────│
|
||||
│ │◀── encrypted response ────────│
|
||||
│ │ │
|
||||
│ streams response back untouched │
|
||||
│◀── encrypted response ────────────────│ │
|
||||
│ │ │
|
||||
SecureClient reads nonce, decrypts │ │
|
||||
SDK SSE processing sees plaintext │ │
|
||||
```
|
||||
|
||||
## Model ID mapping
|
||||
|
||||
Three parties see three different model IDs:
|
||||
|
||||
| Party | Header/Body | Value | Source |
|
||||
|---|---|---|---|
|
||||
| Routstr proxy | `X-Routstr-Model` header | `tinfoil-kimi-k2-6` | SDK sends full caller-facing id |
|
||||
| Tinfoil usage metrics | `model` field | `kimi-k2-6` | Enclave reports the model actually served |
|
||||
| Tinfoil enclave | `body.model` (encrypted) | `kimi-k2-6` | SDK strips `tinfoil-` prefix before encryption |
|
||||
|
||||
## Implementation status
|
||||
|
||||
A dedicated `TinfoilUpstreamProvider` (`routstr/upstream/tinfoil.py`) now
|
||||
implements the direct blind-upstream pattern described above. The shared EHBP
|
||||
helpers in `routstr/upstream/ehbp.py` were extended to:
|
||||
|
||||
- Request usage metrics via `X-Tinfoil-Request-Usage-Metrics: true`.
|
||||
- Parse `X-Tinfoil-Usage-Metrics` from the response header (non-streaming) or
|
||||
HTTP trailer (streaming).
|
||||
- Override the forwarding URL with a validated `X-Tinfoil-Enclave-Url` when the
|
||||
SDK sends it.
|
||||
- Finalize bearer billing with the dedicated EHBP actual-cost finalizer.
|
||||
- Compute X-Cashu refunds from actual cost instead of max cost.
|
||||
|
||||
See `docs/tinfoil-direct-integration.md` for the full implementation notes.
|
||||
|
||||
## Verification status
|
||||
|
||||
Unit coverage includes usage parsing, target validation, HTTP trailer capture,
|
||||
response-size limits, and bearer payment finalization. End-to-end requests have
|
||||
verified both non-streaming usage headers and streaming usage trailers against
|
||||
Tinfoil. SDK behavior on proxy-generated non-2xx responses without an
|
||||
`Ehbp-Response-Nonce` still merits explicit end-to-end coverage.
|
||||
@@ -13,7 +13,10 @@ Before running your node, you should create a `.env` file in the project root. T
|
||||
### Example .env
|
||||
|
||||
```bash
|
||||
ADMIN_PASSWORD=your-secure-password
|
||||
# Encrypts node secrets at rest. Optional — if unset, the node generates a key on
|
||||
# first start and prints it once (back it up). Set it to manage the key yourself
|
||||
# (recommended in production). See "Secrets at Rest" below.
|
||||
ROUTSTR_SECRET_KEY=
|
||||
|
||||
# Node Identity
|
||||
NAME="My AI Node"
|
||||
@@ -25,10 +28,10 @@ RECEIVE_LN_ADDRESS=yourname@wallet.com
|
||||
|
||||
### Setting the UI Password
|
||||
|
||||
There are two ways to set or change your Admin Dashboard password:
|
||||
On first start the node generates an admin password and logs it once — read it from the container logs to sign in. You can then change it two ways:
|
||||
|
||||
1. **Via Environment Variable**: Set `ADMIN_PASSWORD` in your `.env` file before starting the container. This will be the password used for the first login.
|
||||
2. **Via Dashboard**: Once logged in, go to **Settings** → **Security** to update your password. Dashboard settings override the `.env` file once saved.
|
||||
1. **Via Dashboard**: Once logged in, go to **Settings** → **Security** to update your password.
|
||||
2. **Via Environment Variable (legacy seed)**: Setting `ADMIN_PASSWORD` in `.env` before the first start seeds the initial password instead of generating one. It's read only once, for existing deployments; a value left in `.env` is ignored after the node has been configured.
|
||||
|
||||
---
|
||||
|
||||
@@ -123,12 +126,14 @@ Use environment variables for:
|
||||
| -------------------- | --------------------------------- | ------------------------------------ |
|
||||
| `UPSTREAM_BASE_URL` | Upstream API endpoint | — |
|
||||
| `UPSTREAM_API_KEY` | Upstream API key | — |
|
||||
| `ADMIN_PASSWORD` | Dashboard password | (none) |
|
||||
| `ADMIN_PASSWORD` | Legacy seed for the dashboard password (otherwise generated + logged on first start) | (auto-generated) |
|
||||
| `ROUTSTR_SECRET_KEY` | Master key encrypting node secrets at rest. Auto-generated to a key file if unset | (auto-generated) |
|
||||
| `ROUTSTR_SECRET_KEY_FILE` | Path to the generated key file (used when `ROUTSTR_SECRET_KEY` is unset) | `routstr_secret.key` beside the database |
|
||||
| `DATABASE_URL` | Database connection string | `sqlite+aiosqlite:///keys.db` |
|
||||
| `NAME` | Node display name | `ARoutstrNode` |
|
||||
| `DESCRIPTION` | Node description | `A Routstr Node` |
|
||||
| `NPUB` | Nostr public key (bech32) | — |
|
||||
| `NSEC` | Nostr private key | — |
|
||||
| `NSEC` | Legacy seed for the Nostr private key (otherwise set from the admin UI) | — |
|
||||
| `ENABLE_ANALYTICS_SHARING` | Enable usage analytics sharing to Nostr | `true` |
|
||||
| `CASHU_MINTS` | Comma-separated mint URLs | `https://mint.minibits.cash/Bitcoin` |
|
||||
| `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — |
|
||||
@@ -142,6 +147,20 @@ Use environment variables for:
|
||||
|
||||
Environment variables are read on startup. Dashboard settings override them and persist in the database. Once you change a setting in the dashboard, the env var is ignored for that setting.
|
||||
|
||||
### Secrets at Rest
|
||||
|
||||
The node's Nostr private key (`nsec`) is encrypted in the database using
|
||||
`ROUTSTR_SECRET_KEY`. You don't have to set it: if it's unset, the node generates a
|
||||
key on first start, writes it **beside the database** (the file named by
|
||||
`ROUTSTR_SECRET_KEY_FILE`, default `routstr_secret.key`) so it persists on the same
|
||||
volume as your data, and prints it once.
|
||||
|
||||
**Back up that key** — it lives on the same volume as your database, so include it
|
||||
in your backups. If it is lost or changed, previously encrypted secrets can't be
|
||||
decrypted and must be re-entered — there is no rotation. To keep the key off the
|
||||
data volume, set `ROUTSTR_SECRET_KEY` explicitly (an env value always takes
|
||||
precedence over the file). See also [Deployment](deployment.md).
|
||||
|
||||
---
|
||||
|
||||
## Models
|
||||
|
||||
@@ -38,7 +38,6 @@ services:
|
||||
- routstr-data:/app/data
|
||||
environment:
|
||||
DATABASE_URL: "sqlite:////app/data/routstr.db"
|
||||
ADMIN_KEY: "your-secure-admin-key"
|
||||
LOG_LEVEL: "info"
|
||||
|
||||
volumes:
|
||||
@@ -87,6 +86,8 @@ services:
|
||||
- ./logs:/app/logs
|
||||
environment:
|
||||
- TOR_PROXY_URL=socks5://tor:9050
|
||||
# Keep the database (and the key file generated beside it) on the volume.
|
||||
- DATABASE_URL=sqlite:////app/data/routstr.db
|
||||
depends_on:
|
||||
- tor
|
||||
|
||||
@@ -125,8 +126,8 @@ services:
|
||||
- UPSTREAM_BASE_URL=https://api.openai.com/v1
|
||||
- UPSTREAM_API_KEY=sk-proj-...
|
||||
|
||||
# Secure the dashboard (recommended)
|
||||
- ADMIN_PASSWORD=your-secure-password
|
||||
# The admin password is generated and logged once on first start; set
|
||||
# ADMIN_PASSWORD here only as a legacy seed for an existing deployment.
|
||||
|
||||
# Node identity
|
||||
- NAME=My Provider Node
|
||||
@@ -134,6 +135,9 @@ services:
|
||||
|
||||
# Lightning withdrawals
|
||||
- RECEIVE_LN_ADDRESS=me@walletofsatoshi.com
|
||||
|
||||
# Keep the database (and the key file generated beside it) on the volume.
|
||||
- DATABASE_URL=sqlite:////app/data/routstr.db
|
||||
volumes:
|
||||
- ./data:/app/data
|
||||
```
|
||||
@@ -155,22 +159,36 @@ Example `.env`:
|
||||
```bash
|
||||
UPSTREAM_BASE_URL=https://api.openai.com/v1
|
||||
UPSTREAM_API_KEY=sk-proj-...
|
||||
ADMIN_PASSWORD=change-me
|
||||
# Keep the database (and the key file generated beside it) on the mounted volume.
|
||||
DATABASE_URL=sqlite:////app/data/routstr.db
|
||||
# Encrypts node secrets at rest. Optional — if unset, a key is generated next to
|
||||
# your database (on the same volume) and its file is named once for backup. Set
|
||||
# it explicitly to manage the key yourself.
|
||||
ROUTSTR_SECRET_KEY=
|
||||
NAME=My Provider Node
|
||||
RECEIVE_LN_ADDRESS=me@walletofsatoshi.com
|
||||
```
|
||||
|
||||
!!! note "Secret key persistence"
|
||||
If you leave `ROUTSTR_SECRET_KEY` unset, the node generates one and stores it
|
||||
as `routstr_secret.key` **next to your database**, so it persists on the same
|
||||
volume as your data — just include that volume in your backups. For stronger
|
||||
isolation (keeping the key off the data volume), set `ROUTSTR_SECRET_KEY` from
|
||||
a secrets manager instead.
|
||||
|
||||
See [Configuration](configuration.md) for all available options.
|
||||
|
||||
---
|
||||
|
||||
## Persistence
|
||||
|
||||
Routstr stores all data in `/app/data`:
|
||||
Point `DATABASE_URL` inside `/app/data` (as the examples above do) so everything
|
||||
Routstr persists lands on the mounted volume:
|
||||
|
||||
| Path | Contents |
|
||||
|------|----------|
|
||||
| `keys.db` | SQLite database (settings, API keys, sessions) |
|
||||
| `routstr.db` | SQLite database (settings, API keys, sessions) |
|
||||
| `routstr_secret.key` | Auto-generated master key, written beside the database when `ROUTSTR_SECRET_KEY` is unset |
|
||||
| `.wallet/` | Cashu wallet data (your Bitcoin!) |
|
||||
|
||||
!!! warning "Back Up Your Data"
|
||||
|
||||
@@ -29,19 +29,25 @@ In future versions, you'll be able to run a node that connects to other Routstr
|
||||
Create a `.env` file in the root of the project to store your secrets:
|
||||
|
||||
```bash
|
||||
# Initial Admin Password
|
||||
ADMIN_PASSWORD=mysecretpassword
|
||||
# Encrypts node secrets at rest. Optional — if unset, the node generates a key on
|
||||
# first start and prints it once (back it up).
|
||||
ROUTSTR_SECRET_KEY=
|
||||
|
||||
# Node Identity
|
||||
NAME="My AI Node"
|
||||
DESCRIPTION="Fast access to models"
|
||||
NSEC=yournsec
|
||||
|
||||
# Lightning Payouts
|
||||
RECEIVE_LN_ADDRESS=yourname@wallet.com
|
||||
|
||||
```
|
||||
|
||||
The admin password is generated and logged once on first start (read it from the
|
||||
logs to sign in), and your Nostr identity (`nsec`) is configured afterwards from
|
||||
the admin UI — both are stored encrypted in the database, not in `.env`.
|
||||
(`ADMIN_PASSWORD` / `NSEC` are still read once as a legacy seed for existing
|
||||
deployments.)
|
||||
|
||||
## 2. Start the Node
|
||||
|
||||
The recommended way to run Routstr is using Docker Compose, which handles the node, the UI, and optional services like Tor.
|
||||
@@ -72,7 +78,7 @@ docker compose up -d
|
||||
Open the **Admin Dashboard** at [http://localhost:8000/admin/](http://localhost:8000/admin/).
|
||||
|
||||
!!! note "Login"
|
||||
Use the `ADMIN_PASSWORD` you defined in your `.env` file to log in. If you didn't set one, the dashboard will prompt you to set one on first visit.
|
||||
On first start the node generates an admin password and logs it once — read it from the container logs to sign in. You can change it afterwards from **Settings** → **Security**.
|
||||
|
||||
### Connect Your AI Providers
|
||||
|
||||
|
||||
@@ -0,0 +1,493 @@
|
||||
# Tinfoil / PPQ Private-Mode Integration Notes
|
||||
|
||||
This document summarizes the current options for integrating Tinfoil/PPQ private models with Routstr, based on the EHBP work in this branch and local testing against `ppq-private-mode-proxy`.
|
||||
|
||||
## Background
|
||||
|
||||
PPQ private models run behind a Tinfoil/EHBP flow:
|
||||
|
||||
- Request bodies are HPKE-encrypted by a Tinfoil client.
|
||||
- The PPQ `/private/` endpoint routes ciphertext to the attested enclave.
|
||||
- The enclave decrypts, runs inference, and returns an encrypted response.
|
||||
- The caller's Tinfoil client decrypts the response locally.
|
||||
|
||||
The PPQ private-mode proxy (`~/projects/ppq-private-mode-proxy`) uses this pattern with the JavaScript `tinfoil` SDK:
|
||||
|
||||
```ts
|
||||
import { SecureClient } from "tinfoil";
|
||||
|
||||
const apiBase = "https://api.ppq.ai";
|
||||
|
||||
const client = new SecureClient({
|
||||
baseURL: `${apiBase}/private/`,
|
||||
attestationBundleURL: `${apiBase}/private`,
|
||||
transport: "ehbp",
|
||||
});
|
||||
|
||||
await client.ready();
|
||||
|
||||
const response = await client.fetch(`${apiBase}/private/v1/chat/completions`, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
"Content-Type": "application/json",
|
||||
Authorization: `Bearer ${process.env.PPQ_API_KEY}`,
|
||||
"X-Private-Model": "private/gpt-oss-120b",
|
||||
"x-query-source": "api",
|
||||
},
|
||||
body: JSON.stringify({
|
||||
model: "gpt-oss-120b", // enclave-internal model id
|
||||
messages: [{ role: "user", content: "Hello" }],
|
||||
}),
|
||||
});
|
||||
|
||||
const json = await response.json();
|
||||
console.log(json.usage);
|
||||
```
|
||||
|
||||
`X-Private-Model` carries the PPQ-facing private model id, while the encrypted JSON body uses the enclave-internal model id without the `private/` prefix.
|
||||
|
||||
## PPQ private model pricing
|
||||
|
||||
PPQ exposes private model pricing through:
|
||||
|
||||
```text
|
||||
GET https://api.ppq.ai/v1/models?type=all
|
||||
```
|
||||
|
||||
Filter models whose IDs start with `private/`.
|
||||
|
||||
Example private pricing observed:
|
||||
|
||||
| Model | Input USD / 1M tokens | Output USD / 1M tokens |
|
||||
|---|---:|---:|
|
||||
| `private/gpt-oss-120b` | `0.79125` | `1.31875` |
|
||||
| `private/llama3-3-70b` | `1.84625` | `2.90125` |
|
||||
| `private/qwen3-vl-30b` | `1.31875` | `4.22` |
|
||||
| `private/glm-5-2` | `1.5825` | `5.53875` |
|
||||
| `private/gemma4-31b` | `0.47475` | `1.055` |
|
||||
| `private/kimi-k2-6` | `1.5825` | `5.53875` |
|
||||
|
||||
The `ppq-private-mode-proxy` commit `ba984214793d3bca0f7d046b6955d42abc1c6843` changed OpenClaw display metadata to align with a 5% API margin. The live PPQ model endpoint returns the more precise rates above, which already include that margin.
|
||||
|
||||
Actual PPQ private billing is:
|
||||
|
||||
```text
|
||||
price_usd =
|
||||
input_tokens * input_per_1M_tokens / 1_000_000
|
||||
+ output_tokens * output_per_1M_tokens / 1_000_000
|
||||
```
|
||||
|
||||
A test request to `private/gpt-oss-120b` produced:
|
||||
|
||||
```json
|
||||
{
|
||||
"input_count": 74,
|
||||
"output_count": 3,
|
||||
"price_in_usd": 0.00006250875
|
||||
}
|
||||
```
|
||||
|
||||
which exactly matches:
|
||||
|
||||
```text
|
||||
74 * 0.79125 / 1_000_000 + 3 * 1.31875 / 1_000_000
|
||||
= 0.00006250875 USD
|
||||
```
|
||||
|
||||
## Integration architectures
|
||||
|
||||
There are three materially different ways Routstr could integrate Tinfoil/PPQ private inference.
|
||||
|
||||
## Option A: User integrates Tinfoil directly
|
||||
|
||||
```text
|
||||
User app / SDK
|
||||
-> Tinfoil SecureClient / TinfoilAI
|
||||
-> EHBP-encrypted request
|
||||
-> Tinfoil/PPQ private enclave
|
||||
```
|
||||
|
||||
Properties:
|
||||
|
||||
- Best privacy for the user.
|
||||
- Routstr is not in the request path.
|
||||
- User's client encrypts requests and decrypts responses.
|
||||
- Usage is visible to the user's app after decryption.
|
||||
- Billing is handled directly by Tinfoil/PPQ.
|
||||
|
||||
Example with Tinfoil's OpenAI-compatible client:
|
||||
|
||||
```ts
|
||||
import { TinfoilAI } from "tinfoil";
|
||||
|
||||
const client = new TinfoilAI({
|
||||
apiKey: process.env.TINFOIL_API_KEY,
|
||||
transport: "ehbp",
|
||||
});
|
||||
|
||||
const res = await client.chat.completions.create({
|
||||
model: "llama3-3-70b",
|
||||
messages: [{ role: "user", content: "Hello" }],
|
||||
});
|
||||
|
||||
console.log(res.choices[0].message.content);
|
||||
console.log(res.usage);
|
||||
```
|
||||
|
||||
This is not a Routstr marketplace flow unless Routstr only acts as discovery/UI around direct Tinfoil/PPQ usage.
|
||||
|
||||
## Option B: Routstr integrates Tinfoil as an upstream client
|
||||
|
||||
```text
|
||||
User -> Routstr plaintext request
|
||||
-> Routstr Tinfoil SecureClient encrypts to PPQ/Tinfoil
|
||||
-> PPQ private enclave
|
||||
-> Routstr receives decrypted response
|
||||
-> Routstr bills from decrypted usage
|
||||
-> Routstr returns plaintext response to user
|
||||
```
|
||||
|
||||
Properties:
|
||||
|
||||
- Easier exact billing.
|
||||
- Routstr can read the decrypted OpenAI response and `usage` object.
|
||||
- Routstr can charge exact PPQ token pricing.
|
||||
- Privacy is different: the user sends plaintext to Routstr, and Routstr sees prompts/responses.
|
||||
- End-to-end encryption is only Routstr-to-enclave, not user-to-enclave.
|
||||
|
||||
This should be considered a separate product/provider mode, not the same as an end-to-end private relay.
|
||||
|
||||
Practical implementation options:
|
||||
|
||||
1. Run a Node sidecar that uses the `tinfoil` npm package and expose it as a local HTTP upstream to Routstr.
|
||||
2. Port EHBP client behavior to Python.
|
||||
3. Reuse or adapt `ppq-private-mode-proxy` as a local upstream.
|
||||
|
||||
A Node sidecar is probably the quickest implementation path because `ppq-private-mode-proxy` already demonstrates the full flow.
|
||||
|
||||
## Option C: Routstr uses Tinfoil as a direct blind upstream
|
||||
|
||||
This is the current branch's design intent and is likely the best fit if Tinfoil/PPQ exposes usage metadata headers:
|
||||
|
||||
```text
|
||||
User Tinfoil SecureClient
|
||||
-> encrypted request body
|
||||
-> Routstr proxy
|
||||
-> Tinfoil/PPQ private API
|
||||
with X-Tinfoil-Request-Usage-Metrics: true
|
||||
-> PPQ private enclave
|
||||
<- encrypted response
|
||||
plus X-Tinfoil-Usage-Metrics / cost headers
|
||||
<- Routstr proxy
|
||||
-> user decrypts response
|
||||
```
|
||||
|
||||
In this mode Routstr is still a normal upstream proxy from the user's point of view, but the upstream is Tinfoil/PPQ private inference and the body remains opaque to Routstr.
|
||||
|
||||
Properties:
|
||||
|
||||
- Strongest privacy with Routstr in the path.
|
||||
- Routstr never sees plaintext prompt or plaintext response.
|
||||
- Routstr can authenticate and route based on plaintext headers.
|
||||
- Routstr should request Tinfoil usage metadata by adding `X-Tinfoil-Request-Usage-Metrics: true` to the upstream request.
|
||||
- Tinfoil documents `X-Tinfoil-Usage-Metrics` as an upstream response header for non-streaming requests.
|
||||
- For streaming requests, Tinfoil documents `X-Tinfoil-Usage-Metrics` as an HTTP trailer available only after the response body completes.
|
||||
- If Tinfoil/PPQ returns `X-Tinfoil-Usage-Metrics` or a cost header on the encrypted response, Routstr can bill exactly without decrypting the body.
|
||||
- If usage is only present inside the encrypted response body, Routstr still cannot read it and exact billing is not possible without a separate metadata path.
|
||||
|
||||
This is the only architecture that preserves end-to-end encryption from the user to the PPQ/Tinfoil enclave while still letting Routstr mediate payment. The key requirement is that usage/cost metadata must be returned outside the encrypted body, ideally as a response header available before body streaming begins.
|
||||
|
||||
## Current Routstr problem
|
||||
|
||||
The current EHBP implementation charges successful EHBP requests at `max_cost_for_model` because Routstr cannot decrypt the response body:
|
||||
|
||||
```text
|
||||
successful EHBP request -> charge full reserved max cost
|
||||
```
|
||||
|
||||
That is incorrect for PPQ private models. Max cost should be only a reservation/solvency ceiling. Final charge should use actual PPQ private token pricing.
|
||||
|
||||
Desired behavior:
|
||||
|
||||
```text
|
||||
reserve max cost
|
||||
forward encrypted request
|
||||
obtain actual usage/cost metadata
|
||||
finalize actual cost
|
||||
refund/release the difference
|
||||
```
|
||||
|
||||
## Usage/cost metadata requirement
|
||||
|
||||
For blind-relay exact billing, PPQ should return one of the following outside the encrypted response body:
|
||||
|
||||
```http
|
||||
X-PPQ-Cost-USD: 0.00006250875
|
||||
```
|
||||
|
||||
or:
|
||||
|
||||
```http
|
||||
X-Private-Usage-Metrics: input=74,output=3
|
||||
```
|
||||
|
||||
or:
|
||||
|
||||
```http
|
||||
X-Tinfoil-Usage-Metrics: prompt=74,completion=3,total=77
|
||||
```
|
||||
|
||||
Tinfoil proxy documentation references a usage-metrics flow where a proxy can request usage via:
|
||||
|
||||
```http
|
||||
X-Tinfoil-Request-Usage-Metrics: true
|
||||
```
|
||||
|
||||
and read usage from:
|
||||
|
||||
```http
|
||||
X-Tinfoil-Usage-Metrics
|
||||
```
|
||||
|
||||
Docs/example references:
|
||||
|
||||
- https://docs.tinfoil.sh/guides/proxy-server
|
||||
- https://github.com/tinfoilsh/encrypted-request-proxy-example
|
||||
|
||||
During local PPQ testing, PPQ responses included this CORS exposure header:
|
||||
|
||||
```http
|
||||
Access-Control-Expose-Headers: Ehbp-Response-Nonce, X-Private-Usage-Metrics, X-Encrypted-Usage-Metrics, X-Tinfoil-Usage-Metrics
|
||||
```
|
||||
|
||||
However, the actual tested non-streaming response did not include any of these usage headers, even when `X-Tinfoil-Request-Usage-Metrics: true` was sent.
|
||||
|
||||
The decrypted body did include normal OpenAI usage, but only the decrypting Tinfoil client can see that body.
|
||||
|
||||
## Query-history fallback
|
||||
|
||||
PPQ's query history endpoint exposes actual usage and cost:
|
||||
|
||||
```text
|
||||
GET https://api.ppq.ai/queries/history?page=1&page_count=...
|
||||
```
|
||||
|
||||
A record includes:
|
||||
|
||||
```json
|
||||
{
|
||||
"timestamp": "...",
|
||||
"model": "private/gpt-oss-120b",
|
||||
"input_count": 74,
|
||||
"output_count": 3,
|
||||
"price_in_usd": 0.00006250875,
|
||||
"query_type": "chat_completion",
|
||||
"query_source": "api"
|
||||
}
|
||||
```
|
||||
|
||||
This could be used as a fallback, but it is less robust than response headers/trailers because matching a request to a history row can be race-prone under concurrency. It would need a reliable request identifier or metadata field that PPQ stores in history.
|
||||
|
||||
## Recommended Routstr direction
|
||||
|
||||
Use Tinfoil/PPQ as a direct blind upstream and have Routstr explicitly request usage metadata:
|
||||
|
||||
```text
|
||||
User encrypts body
|
||||
Routstr reserves max cost
|
||||
Routstr forwards encrypted body to Tinfoil/PPQ /private/
|
||||
Routstr includes X-Tinfoil-Request-Usage-Metrics: true
|
||||
Tinfoil/PPQ returns X-Tinfoil-Usage-Metrics as a response header for non-streaming,
|
||||
or as an HTTP trailer after the body completes for streaming
|
||||
Routstr finalizes exact charge
|
||||
User decrypts encrypted response
|
||||
```
|
||||
|
||||
This preserves both:
|
||||
|
||||
- privacy: Routstr cannot read prompts/responses;
|
||||
- exact billing: Routstr can charge actual PPQ private model cost, assuming Tinfoil/PPQ returns usage/cost metadata outside the encrypted body.
|
||||
|
||||
Implementation steps:
|
||||
|
||||
1. Update PPQ model fetching to include private models:
|
||||
|
||||
```text
|
||||
GET https://api.ppq.ai/v1/models?type=all
|
||||
```
|
||||
|
||||
2. Register `private/*` models and any Routstr-facing aliases with correct `forwarded_model_id`.
|
||||
|
||||
3. Keep max-cost reservation for bearer keys and X-Cashu solvency checks.
|
||||
|
||||
4. Add parsing support for possible usage/cost headers:
|
||||
|
||||
```http
|
||||
X-PPQ-Cost-USD
|
||||
X-Private-Usage-Metrics
|
||||
X-Tinfoil-Usage-Metrics
|
||||
X-Encrypted-Usage-Metrics
|
||||
```
|
||||
|
||||
5. Finalize by actual cost instead of max cost:
|
||||
|
||||
```text
|
||||
actual_msats = ceil((actual_usd / sats_usd_price()) * 1000)
|
||||
actual_msats = max(actual_msats, settings.min_request_msat)
|
||||
actual_msats = min(actual_msats, reserved_msats)
|
||||
```
|
||||
|
||||
or, if only token counts are available:
|
||||
|
||||
```text
|
||||
actual_usd =
|
||||
input_tokens * input_per_1M_tokens / 1_000_000
|
||||
+ output_tokens * output_per_1M_tokens / 1_000_000
|
||||
```
|
||||
|
||||
6. If usage/cost metadata is missing, choose an explicit policy:
|
||||
|
||||
- fail closed and refund/revert;
|
||||
- query PPQ history as a fallback;
|
||||
- fallback to max-cost billing only if explicitly configured and clearly disclosed.
|
||||
|
||||
Silent max-cost billing should not be the default for PPQ private requests.
|
||||
|
||||
## X-Cashu consideration
|
||||
|
||||
For bearer-auth requests, finalization can happen after the response stream completes if usage is delivered as a trailer.
|
||||
|
||||
For `X-Cashu`, Routstr needs to return the refund token in the response headers. If usage/cost is only available after consuming the encrypted response stream, Routstr may need to buffer EHBP responses before sending them to the client so it can compute the refund amount first.
|
||||
|
||||
Possible approaches:
|
||||
|
||||
1. Prefer a non-trailer response header with actual cost, available before streaming body starts.
|
||||
2. Buffer EHBP X-Cashu responses and then return `X-Cashu` refund.
|
||||
3. Introduce a later/refund-claim mechanism, which would be a larger protocol change.
|
||||
|
||||
## Summary
|
||||
|
||||
- PPQ private models are billed per actual input/output tokens.
|
||||
- Private model rates are available from `GET /v1/models?type=all`.
|
||||
- Current Routstr EHBP billing at max cost is wrong for PPQ private models.
|
||||
- Direct Tinfoil integration inside Routstr would enable exact usage billing but would make Routstr see plaintext.
|
||||
- A blind EHBP relay preserves privacy but requires PPQ/Tinfoil to expose usage/cost in plaintext headers/trailers.
|
||||
- The preferred solution is to keep Routstr blind and have PPQ return billing metadata outside the encrypted body.
|
||||
|
||||
## Implementation status
|
||||
|
||||
Direct Tinfoil upstream integration is implemented in `routstr/upstream/tinfoil.py`
|
||||
and `routstr/upstream/ehbp.py`.
|
||||
|
||||
### What was built
|
||||
|
||||
- `TinfoilUpstreamProvider` (`provider_type = "tinfoil"`):
|
||||
- Base URL: `https://inference.tinfoil.sh`
|
||||
- Fetches models from the public `GET /v1/models` endpoint (no auth needed).
|
||||
- Parses Tinfoil's pricing (`inputTokenPricePer1M`, `outputTokenPricePer1M`,
|
||||
`requestPrice`) into the standard `Model`/`Pricing` schema.
|
||||
- `supports_ehbp = True` — acts as a blind EHBP relay.
|
||||
- `get_ehbp_forwarding_target()` returns a target that includes
|
||||
`X-Tinfoil-Request-Usage-Metrics: true`.
|
||||
- `forward_get_request()` proxies `/attestation` to `https://atc.tinfoil.sh/attestation`
|
||||
so the SDK can fetch attestation bundles through Routstr.
|
||||
- Registered in `routstr/upstream/__init__.py` and seeded from
|
||||
`TINFOIL_API_KEY` env var.
|
||||
|
||||
- `routstr/upstream/ehbp.py`:
|
||||
- `parse_tinfoil_usage_metrics()` parses
|
||||
`prompt=N,completion=N[,total=N][,model=<name>]` into an OpenAI-style
|
||||
usage dict. The `model` field (added in tinfoilsh/confidential-model-router
|
||||
PR #385) is extracted as a string.
|
||||
- `_resolve_ehbp_target_url()` overrides the forwarding URL with
|
||||
`X-Tinfoil-Enclave-Url` when the SDK sends it.
|
||||
- `_strip_proxy_headers()` removes `X-Routstr-Model`,
|
||||
`X-Tinfoil-Enclave-Url`, and `X-Tinfoil-Request-Usage-Metrics` before
|
||||
forwarding to the enclave.
|
||||
- `_compute_ehbp_actual_cost()` converts the usage header into msats via
|
||||
`calculate_cost()`, clamped to `[min_request_msat, max_cost_for_model]`.
|
||||
When the header's `model=<name>` differs from the requested model, the
|
||||
actual served model's pricing is used for cost calculation.
|
||||
- `forward_ehbp_request()` (bearer auth): if `X-Tinfoil-Usage-Metrics` is
|
||||
present in the response header, finalizes with `adjust_payment_for_tokens()`
|
||||
for exact billing; otherwise falls back to max-cost. Billing uses the
|
||||
actual served model when it differs from the requested one.
|
||||
- `forward_ehbp_x_cashu_request()`: if usage is available, computes the
|
||||
refund from actual cost instead of max cost, using the actual served
|
||||
model's pricing when applicable.
|
||||
|
||||
- `routstr/proxy.py`: `/attestation` and `/tee/attestation` paths are forwarded
|
||||
to Tinfoil upstreams without model/cost/auth lookups.
|
||||
|
||||
### Billing behavior
|
||||
|
||||
| Request shape | Usage source | Billing |
|
||||
|---|---|---|
|
||||
| Bearer, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Exact token cost via `adjust_payment_for_tokens` |
|
||||
| Bearer, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Exact token cost (h11 captures trailers) |
|
||||
| Bearer, no usage header/trailer | N/A | Max-cost fallback |
|
||||
| X-Cashu, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Refund = `redeemed - actual_cost` |
|
||||
| X-Cashu, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Refund = `redeemed - actual_cost` (h11 captures trailers) |
|
||||
| X-Cashu, no usage header/trailer | N/A | Refund = `redeemed - max_cost` |
|
||||
|
||||
### Cost response headers
|
||||
|
||||
Since EHBP response bodies are opaque encrypted blobs, per-request cost cannot
|
||||
be injected into the JSON body (as done in the normal proxy flow). Instead,
|
||||
Routstr returns cost info as response headers:
|
||||
|
||||
| Header | Auth | Description |
|
||||
|---|---|---|
|
||||
| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Total msats charged for this request |
|
||||
| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the charge |
|
||||
| `X-Routstr-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens |
|
||||
| `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens |
|
||||
|
||||
The client/Tinfoil SDK can read these headers from the HTTP response without
|
||||
needing to decrypt the body.
|
||||
|
||||
### Setup
|
||||
|
||||
```bash
|
||||
TINFOIL_API_KEY=your-tinfoil-api-key
|
||||
```
|
||||
|
||||
The provider is auto-seeded on first startup.
|
||||
|
||||
### Usage metrics header format
|
||||
|
||||
Tinfoil returns usage metrics in the `X-Tinfoil-Usage-Metrics` response header
|
||||
(non-streaming) or HTTP trailer (streaming) when `X-Tinfoil-Request-Usage-Metrics:
|
||||
true` is sent. As of tinfoilsh/confidential-model-router PR #385, the format is:
|
||||
|
||||
```
|
||||
prompt=<prompt_tokens>,completion=<completion_tokens>,total=<total_tokens>,model=<served_model>
|
||||
```
|
||||
|
||||
The `model` field carries the actual model name served by the enclave.
|
||||
Routstr uses this to:
|
||||
|
||||
- Verify the served model matches the expected upstream model. The comparison
|
||||
uses ``model_obj.forwarded_model_id`` (the actual upstream ID, e.g.
|
||||
``glm-5-2``) rather than ``model_obj.id`` (the client-facing alias, e.g.
|
||||
``tinfoil-glm-5-2``), so aliased models don't trigger a spurious mismatch.
|
||||
- When they genuinely differ (Tinfoil served a different upstream model than
|
||||
expected), look up the actual served model's pricing and use it for billing.
|
||||
The reverse lookup uses ``get_model_instance``, which resolves
|
||||
``forwarded_model_id`` values registered as routable aliases.
|
||||
- Log the discrepancy for observability.
|
||||
|
||||
If the actual model is not found in Routstr's model registry, billing falls
|
||||
back to the requested model's pricing.
|
||||
|
||||
### What still needs verification
|
||||
|
||||
- ~~End-to-end test with a real Tinfoil SDK client against a Routstr node with
|
||||
`TINFOIL_API_KEY` set.~~ Verified: both non-streaming (header) and streaming
|
||||
(trailer) responses include `model=<name>`.
|
||||
- Streaming trailer capture is implemented by buffering the encrypted response
|
||||
in `forward_with_trailer()` and then using the dedicated EHBP payment
|
||||
finalizers for bearer and X-Cashu requests. This provides actual-cost billing
|
||||
today, at the cost of full time-to-last-byte latency for streaming responses.
|
||||
- Whether Tinfoil's `/v1/responses` endpoint also returns usage metrics
|
||||
headers or trailers.
|
||||
@@ -0,0 +1,37 @@
|
||||
"""add fee payout checkpoint
|
||||
|
||||
Revision ID: d7e8f9a0b1c2
|
||||
Revises: c6d7e8f9a0b1
|
||||
Create Date: 2026-07-18 00:00:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "d7e8f9a0b1c2"
|
||||
down_revision = "c6d7e8f9a0b1"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"routstr_fees",
|
||||
sa.Column(
|
||||
"payout_in_progress_msats",
|
||||
sa.Integer(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
),
|
||||
)
|
||||
op.add_column(
|
||||
"routstr_fees",
|
||||
sa.Column("payout_started_at", sa.Integer(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("routstr_fees", "payout_started_at")
|
||||
op.drop_column("routstr_fees", "payout_in_progress_msats")
|
||||
@@ -0,0 +1,51 @@
|
||||
"""add secrets table
|
||||
|
||||
Revision ID: fc4fa29630d2
|
||||
Revises: d7e8f9a0b1c2
|
||||
Create Date: 2026-07-23 00:00:00.000000
|
||||
|
||||
Creates the node-level singleton secret store (issue #553). Schema only; moving
|
||||
any legacy plaintext into the encrypted/hashed columns happens at bootstrap,
|
||||
where the live ROUTSTR_SECRET_KEY is available. ``nsec_state`` records the vault's
|
||||
ownership of the nsec (legacy | encrypted | cleared), so a cleared identity is
|
||||
never resurrected from a stale legacy ``NSEC`` env var / settings blob on the next
|
||||
boot.
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
revision = "fc4fa29630d2"
|
||||
down_revision = "d7e8f9a0b1c2"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"secrets",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column(
|
||||
"admin_password_hash",
|
||||
sqlmodel.sql.sqltypes.AutoString(),
|
||||
nullable=True,
|
||||
),
|
||||
sa.Column(
|
||||
"encrypted_nsec",
|
||||
sqlmodel.sql.sqltypes.AutoString(),
|
||||
nullable=True,
|
||||
),
|
||||
sa.Column(
|
||||
"nsec_state",
|
||||
sqlmodel.sql.sqltypes.AutoString(),
|
||||
nullable=False,
|
||||
server_default="legacy",
|
||||
),
|
||||
sa.Column("updated_at", sa.Integer(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("secrets")
|
||||
@@ -10,6 +10,7 @@ dependencies = [
|
||||
"aiosqlite>=0.20",
|
||||
"sqlmodel>=0.0.24",
|
||||
"httpx[socks]>=0.25.2",
|
||||
"h11>=0.14",
|
||||
"greenlet>=3.2.1",
|
||||
"alembic>=1.13",
|
||||
"python-json-logger>=2.0.0",
|
||||
|
||||
+30
-20
@@ -86,10 +86,12 @@ def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
|
||||
|
||||
def create_model_mappings(
|
||||
upstreams: list["BaseUpstreamProvider"],
|
||||
overrides_by_id: dict[str, tuple],
|
||||
disabled_model_ids: set[str],
|
||||
overrides_by_key: dict[tuple[str, int], tuple],
|
||||
disabled_model_keys: set[tuple[str, int]],
|
||||
) -> tuple[
|
||||
dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"]
|
||||
dict[str, "Model"],
|
||||
dict[str, list[tuple["Model", "BaseUpstreamProvider"]]],
|
||||
dict[str, "Model"],
|
||||
]:
|
||||
"""Create optimal model mappings based on cost and provider preferences.
|
||||
|
||||
@@ -97,7 +99,9 @@ def create_model_mappings(
|
||||
and creates three mappings based on cost optimization:
|
||||
|
||||
1. model_instances: alias -> Model (all model aliases mapped to their Model objects)
|
||||
2. provider_map: alias -> List[UpstreamProvider] (sorted list of providers for each alias)
|
||||
2. provider_map: alias -> List[(Model, UpstreamProvider)] (sorted candidate
|
||||
list for each alias; each provider is paired with ITS OWN model so
|
||||
failover can forward and bill the candidate that actually serves)
|
||||
3. unique_models: base_id -> Model (unique models without provider prefixes)
|
||||
|
||||
The algorithm:
|
||||
@@ -107,8 +111,9 @@ def create_model_mappings(
|
||||
|
||||
Args:
|
||||
upstreams: List of all upstream provider instances
|
||||
overrides_by_id: Dict of model overrides from database {model_id: (ModelRow, fee)}
|
||||
disabled_model_ids: Set of model IDs that should be excluded
|
||||
overrides_by_key: Dict of model overrides from database
|
||||
{(model_id_lower, upstream_provider_id): (ModelRow, fee)}
|
||||
disabled_model_keys: Set of provider-scoped model keys that should be excluded
|
||||
|
||||
Returns:
|
||||
Tuple of (model_instances, provider_map, unique_models)
|
||||
@@ -179,14 +184,22 @@ def create_model_mappings(
|
||||
"""Process all models from a given provider."""
|
||||
upstream_prefix = getattr(upstream, "upstream_name", None)
|
||||
provider_key = get_provider_identity(upstream)
|
||||
upstream_db_id = getattr(upstream, "db_id", None)
|
||||
|
||||
for model in upstream.get_cached_models():
|
||||
if not model.enabled or model.id in disabled_model_ids:
|
||||
model_key = (
|
||||
(model.id.lower(), upstream_db_id)
|
||||
if isinstance(upstream_db_id, int)
|
||||
else None
|
||||
)
|
||||
if not model.enabled or (
|
||||
model_key is not None and model_key in disabled_model_keys
|
||||
):
|
||||
continue
|
||||
|
||||
# Apply overrides if present
|
||||
if model.id in overrides_by_id:
|
||||
override_row, provider_fee = overrides_by_id[model.id]
|
||||
# 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
|
||||
)
|
||||
@@ -237,13 +250,10 @@ def create_model_mappings(
|
||||
|
||||
# Include enabled DB overrides even when provider discovery misses models.
|
||||
# This is important for deployment-based providers like Azure.
|
||||
for model_id, override_data in overrides_by_id.items():
|
||||
if model_id in disabled_model_ids:
|
||||
for (model_id, upstream_provider_id), override_data in overrides_by_key.items():
|
||||
if (model_id, upstream_provider_id) in disabled_model_keys:
|
||||
continue
|
||||
override_row, provider_fee = override_data
|
||||
upstream_provider_id = getattr(override_row, "upstream_provider_id", None)
|
||||
if not isinstance(upstream_provider_id, int):
|
||||
continue
|
||||
|
||||
upstream_for_override = providers_by_db_id.get(upstream_provider_id)
|
||||
if upstream_for_override is None:
|
||||
@@ -321,7 +331,7 @@ def create_model_mappings(
|
||||
|
||||
# Sort candidates and build final maps
|
||||
model_instances: dict[str, "Model"] = {}
|
||||
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
|
||||
provider_map: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
|
||||
|
||||
def alias_priority(model: "Model", alias: str) -> int:
|
||||
"""Rank how strong the mapping of alias->model is.
|
||||
@@ -368,13 +378,13 @@ def create_model_mappings(
|
||||
|
||||
best_model, best_provider = items[0]
|
||||
model_instances[alias] = best_model
|
||||
provider_map[alias] = [p for _, p in items]
|
||||
provider_map[alias] = list(items)
|
||||
|
||||
# Log provider distribution (using top provider for stats)
|
||||
provider_counts: dict[str, int] = {}
|
||||
for providers in provider_map.values():
|
||||
if providers:
|
||||
provider = providers[0]
|
||||
for candidate_list in provider_map.values():
|
||||
if candidate_list:
|
||||
provider = candidate_list[0][1]
|
||||
provider_name = getattr(provider, "upstream_name", "unknown")
|
||||
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
|
||||
|
||||
|
||||
+13
-2
@@ -4,7 +4,7 @@ import math
|
||||
import random
|
||||
import time
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import case
|
||||
@@ -26,6 +26,9 @@ from .wallet import (
|
||||
deserialize_token_from_string,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .payment.models import Model
|
||||
|
||||
logger = get_logger(__name__)
|
||||
payments_logger = get_logger("routstr.payments")
|
||||
|
||||
@@ -771,12 +774,18 @@ async def adjust_payment_for_tokens(
|
||||
response_data: dict,
|
||||
session: AsyncSession,
|
||||
deducted_max_cost: int,
|
||||
model_obj: "Model | None",
|
||||
provider_fee: float | None,
|
||||
) -> dict:
|
||||
"""
|
||||
Adjusts the payment based on token usage in the response.
|
||||
This is called after the initial payment and the upstream request is complete.
|
||||
Returns cost data to be included in the response.
|
||||
|
||||
``model_obj`` is the model that actually served the request; it is passed
|
||||
through to ``calculate_cost`` so billing uses the serving candidate's
|
||||
pricing instead of re-deriving it from the response's model string.
|
||||
|
||||
The response's usage object is normalized with the default union parser in
|
||||
``calculate_cost``.
|
||||
"""
|
||||
@@ -861,7 +870,9 @@ async def adjust_payment_for_tokens(
|
||||
extra={"error": str(e), "fee_msats": fee_msats},
|
||||
)
|
||||
|
||||
match await calculate_cost(response_data, deducted_max_cost):
|
||||
match await calculate_cost(
|
||||
response_data, deducted_max_cost, model_obj, provider_fee
|
||||
):
|
||||
case MaxCostData() as cost:
|
||||
logger.debug(
|
||||
"Using max cost data (no token adjustment)",
|
||||
|
||||
+3
-1
@@ -15,7 +15,9 @@ from .core.db import (
|
||||
AsyncSession,
|
||||
CashuTransaction,
|
||||
get_session,
|
||||
store_cashu_transaction,
|
||||
)
|
||||
from .core.db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from .core.logging import get_logger
|
||||
from .core.settings import settings
|
||||
|
||||
+90
-46
@@ -20,6 +20,7 @@ from ..wallet import (
|
||||
send_token,
|
||||
slow_filter_spend_proofs,
|
||||
)
|
||||
from . import vault
|
||||
from .db import (
|
||||
ApiKey,
|
||||
CashuTransaction,
|
||||
@@ -28,11 +29,17 @@ from .db import (
|
||||
ModelRow,
|
||||
UpstreamProviderRow,
|
||||
create_session,
|
||||
get_secret,
|
||||
set_admin_password,
|
||||
set_nsec,
|
||||
)
|
||||
from .db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from .log_manager import log_manager
|
||||
from .logging import get_logger
|
||||
from .provider_slugs import allocate_unique_provider_slug
|
||||
from .settings import SettingsService, settings
|
||||
from .settings import SettingsService, derive_npub_from_nsec, settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -203,8 +210,6 @@ async def get_settings(request: Request) -> dict:
|
||||
data = settings.dict()
|
||||
if "upstream_api_key" in data:
|
||||
data["upstream_api_key"] = "[REDACTED]" if data["upstream_api_key"] else ""
|
||||
if "admin_password" in data:
|
||||
data["admin_password"] = "[REDACTED]" if data["admin_password"] else ""
|
||||
if "nsec" in data:
|
||||
data["nsec"] = "[REDACTED]" if data["nsec"] else ""
|
||||
return data
|
||||
@@ -221,9 +226,10 @@ class PasswordUpdate(BaseModel):
|
||||
|
||||
@admin_router.patch("/api/settings", dependencies=[Depends(require_admin_api)])
|
||||
async def update_settings(request: Request, update: SettingsUpdate) -> dict:
|
||||
# Remove sensitive fields from general settings update
|
||||
# Secrets are not editable through the general settings endpoint; they have
|
||||
# dedicated rotation paths and never reach the settings blob.
|
||||
settings_data = update.root.copy()
|
||||
sensitive_fields = ["admin_password", "upstream_api_key", "nsec"]
|
||||
sensitive_fields = ["upstream_api_key", "nsec"]
|
||||
for field in sensitive_fields:
|
||||
if field in settings_data:
|
||||
del settings_data[field]
|
||||
@@ -238,8 +244,6 @@ async def update_settings(request: Request, update: SettingsUpdate) -> dict:
|
||||
data = new_settings.dict()
|
||||
if "upstream_api_key" in data:
|
||||
data["upstream_api_key"] = "[REDACTED]" if data["upstream_api_key"] else ""
|
||||
if "admin_password" in data:
|
||||
data["admin_password"] = "[REDACTED]" if data["admin_password"] else ""
|
||||
if "nsec" in data:
|
||||
data["nsec"] = "[REDACTED]" if data["nsec"] else ""
|
||||
return data
|
||||
@@ -247,44 +251,63 @@ async def update_settings(request: Request, update: SettingsUpdate) -> dict:
|
||||
|
||||
@admin_router.patch("/api/password", dependencies=[Depends(require_admin_api)])
|
||||
async def update_password(request: Request, password_update: PasswordUpdate) -> dict:
|
||||
current_password = settings.admin_password
|
||||
|
||||
if not current_password:
|
||||
raise HTTPException(status_code=500, detail="Admin password not configured")
|
||||
|
||||
if password_update.current_password != current_password:
|
||||
raise HTTPException(status_code=401, detail="Current password is incorrect")
|
||||
|
||||
# Validate new password
|
||||
new_password = password_update.new_password.strip()
|
||||
if len(new_password) < 6:
|
||||
raise HTTPException(
|
||||
status_code=400, detail="New password must be at least 6 characters"
|
||||
)
|
||||
|
||||
# Update password
|
||||
async with create_session() as session:
|
||||
await SettingsService.update({"admin_password": new_password}, session)
|
||||
secret = await get_secret(session)
|
||||
|
||||
if not secret.admin_password_hash:
|
||||
raise HTTPException(
|
||||
status_code=500, detail="Admin password not configured"
|
||||
)
|
||||
|
||||
if not vault.verify_password(
|
||||
password_update.current_password, secret.admin_password_hash
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=401, detail="Current password is incorrect"
|
||||
)
|
||||
|
||||
# Validate new password
|
||||
new_password = password_update.new_password.strip()
|
||||
if len(new_password) < vault.MIN_PASSWORD_LENGTH:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"New password must be at least "
|
||||
f"{vault.MIN_PASSWORD_LENGTH} characters"
|
||||
),
|
||||
)
|
||||
|
||||
await set_admin_password(session, new_password)
|
||||
|
||||
return {"ok": True, "message": "Password updated successfully"}
|
||||
|
||||
|
||||
class SetupRequest(BaseModel):
|
||||
password: str
|
||||
class NsecUpdate(BaseModel):
|
||||
nsec: str
|
||||
|
||||
|
||||
@admin_router.post("/api/setup")
|
||||
async def initial_setup(request: Request, payload: SetupRequest) -> dict[str, object]:
|
||||
if settings.admin_password:
|
||||
raise HTTPException(status_code=409, detail="Admin password already set")
|
||||
pw = (payload.password or "").strip()
|
||||
if len(pw) < 8:
|
||||
raise HTTPException(
|
||||
status_code=400, detail="Password must be at least 8 characters"
|
||||
)
|
||||
@admin_router.patch("/api/nsec", dependencies=[Depends(require_admin_api)])
|
||||
async def update_nsec(request: Request, payload: NsecUpdate) -> dict[str, object]:
|
||||
# The node's Nostr identity is a secret: it is stored encrypted in the
|
||||
# Secret store, never in the settings blob, so it gets its own endpoint
|
||||
# rather than riding the general settings PATCH (which strips it). An empty
|
||||
# nsec clears the identity.
|
||||
nsec = payload.nsec.strip()
|
||||
npub = ""
|
||||
if nsec:
|
||||
derived = derive_npub_from_nsec(nsec)
|
||||
if not derived:
|
||||
raise HTTPException(status_code=400, detail="Invalid nsec")
|
||||
npub = derived
|
||||
|
||||
async with create_session() as session:
|
||||
await SettingsService.update({"admin_password": pw}, session)
|
||||
return {"ok": True}
|
||||
await set_nsec(session, nsec)
|
||||
|
||||
# Reflect the change in the live runtime so Nostr signing/announcements pick
|
||||
# it up without a restart (mirrors what bootstrap_secrets sets at boot).
|
||||
settings.nsec = nsec
|
||||
settings.npub = npub
|
||||
return {"ok": True, "npub": npub}
|
||||
|
||||
|
||||
class AdminLoginRequest(BaseModel):
|
||||
@@ -295,12 +318,16 @@ class AdminLoginRequest(BaseModel):
|
||||
async def admin_login(
|
||||
request: Request, payload: AdminLoginRequest
|
||||
) -> dict[str, object]:
|
||||
admin_pw = settings.admin_password
|
||||
async with create_session() as session:
|
||||
secret = await get_secret(session)
|
||||
# Read the hash while the session is open; the ORM object is detached
|
||||
# once the context exits and its attributes can no longer be loaded.
|
||||
password_hash = secret.admin_password_hash
|
||||
|
||||
if not admin_pw:
|
||||
if not password_hash:
|
||||
raise HTTPException(status_code=500, detail="Admin password not configured")
|
||||
|
||||
if payload.password != admin_pw:
|
||||
if not vault.verify_password(payload.password, password_hash):
|
||||
raise HTTPException(status_code=401, detail="Invalid password")
|
||||
|
||||
token = secrets.token_urlsafe(32)
|
||||
@@ -411,12 +438,11 @@ async def withdraw(
|
||||
# Get wallet and check balance
|
||||
from .settings import settings as global_settings
|
||||
|
||||
wallet = await get_wallet(
|
||||
withdraw_request.mint_url or global_settings.primary_mint, withdraw_request.unit
|
||||
)
|
||||
effective_mint = withdraw_request.mint_url or global_settings.primary_mint
|
||||
wallet = await get_wallet(effective_mint, withdraw_request.unit)
|
||||
proofs = get_proofs_per_mint_and_unit(
|
||||
wallet,
|
||||
withdraw_request.mint_url or global_settings.primary_mint,
|
||||
effective_mint,
|
||||
withdraw_request.unit,
|
||||
not_reserved=True,
|
||||
)
|
||||
@@ -432,8 +458,27 @@ async def withdraw(
|
||||
raise HTTPException(status_code=400, detail="Insufficient wallet balance")
|
||||
|
||||
token = await send_token(
|
||||
withdraw_request.amount, withdraw_request.unit, withdraw_request.mint_url
|
||||
withdraw_request.amount, withdraw_request.unit, effective_mint
|
||||
)
|
||||
try:
|
||||
await store_cashu_transaction(
|
||||
token=token,
|
||||
amount=withdraw_request.amount,
|
||||
unit=withdraw_request.unit,
|
||||
mint_url=effective_mint,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="admin",
|
||||
)
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Admin withdrawal token issued without a persisted audit record",
|
||||
extra={
|
||||
"amount": withdraw_request.amount,
|
||||
"unit": withdraw_request.unit,
|
||||
"mint_url": effective_mint,
|
||||
},
|
||||
)
|
||||
return {"token": token}
|
||||
|
||||
|
||||
@@ -461,7 +506,6 @@ class ModelCreate(BaseModel):
|
||||
async def upsert_provider_model(
|
||||
provider_id: str, payload: ModelCreate
|
||||
) -> dict[str, object]:
|
||||
print(payload)
|
||||
logger.info(
|
||||
f"UPSERT_PROVIDER_MODEL called: provider_id={provider_id}, model_id={payload.id}"
|
||||
)
|
||||
|
||||
+208
-9
@@ -1,16 +1,19 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
import pathlib
|
||||
import sqlite3
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from enum import Enum
|
||||
from typing import AsyncGenerator
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from alembic.util.exc import CommandError
|
||||
from sqlalchemy import UniqueConstraint, delete
|
||||
from sqlalchemy.exc import OperationalError
|
||||
from sqlalchemy.exc import IntegrityError, OperationalError
|
||||
from sqlalchemy.ext.asyncio.engine import create_async_engine
|
||||
from sqlalchemy.orm import aliased
|
||||
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
|
||||
@@ -287,10 +290,13 @@ async def store_cashu_transaction(
|
||||
created_at: int | None = None,
|
||||
source: str = "x-cashu",
|
||||
api_key_hashed_key: str | None = None,
|
||||
) -> None:
|
||||
transaction_id: str | None = None,
|
||||
log_failure: bool = True,
|
||||
) -> bool:
|
||||
try:
|
||||
async with create_session() as session:
|
||||
tx = CashuTransaction(
|
||||
id=transaction_id or uuid.uuid4().hex,
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit=unit,
|
||||
@@ -304,11 +310,93 @@ async def store_cashu_transaction(
|
||||
)
|
||||
session.add(tx)
|
||||
await session.commit()
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to store cashu transaction: {e} (type={typ})",
|
||||
extra={"error": str(e), "type": typ},
|
||||
)
|
||||
except Exception:
|
||||
if log_failure:
|
||||
logger.critical(
|
||||
"Failed to store Cashu transaction",
|
||||
extra={"type": typ, "request_id": request_id, "source": source},
|
||||
exc_info=True,
|
||||
)
|
||||
raise
|
||||
return True
|
||||
|
||||
|
||||
async def _cashu_transaction_exists(transaction_id: str) -> bool:
|
||||
async with create_session() as session:
|
||||
return await session.get(CashuTransaction, transaction_id) is not None
|
||||
|
||||
|
||||
async def store_cashu_transaction_with_retry(
|
||||
token: str,
|
||||
amount: int,
|
||||
unit: str,
|
||||
mint_url: str | None = None,
|
||||
typ: str = "out",
|
||||
request_id: str | None = None,
|
||||
collected: bool = False,
|
||||
created_at: int | None = None,
|
||||
source: str = "x-cashu",
|
||||
api_key_hashed_key: str | None = None,
|
||||
max_attempts: int = 3,
|
||||
) -> bool:
|
||||
"""Retry a critical Cashu transaction write with bounded backoff."""
|
||||
transaction_id = hashlib.sha256(f"{typ}\0{token}".encode()).hexdigest()
|
||||
last_error: Exception | None = None
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
try:
|
||||
return await store_cashu_transaction(
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit=unit,
|
||||
mint_url=mint_url,
|
||||
typ=typ,
|
||||
request_id=request_id,
|
||||
collected=collected,
|
||||
created_at=created_at,
|
||||
source=source,
|
||||
api_key_hashed_key=api_key_hashed_key,
|
||||
transaction_id=transaction_id,
|
||||
log_failure=False,
|
||||
)
|
||||
except IntegrityError as error:
|
||||
try:
|
||||
if await _cashu_transaction_exists(transaction_id):
|
||||
return True
|
||||
except Exception as lookup_error:
|
||||
last_error = lookup_error
|
||||
else:
|
||||
last_error = error
|
||||
except Exception as error:
|
||||
last_error = error
|
||||
|
||||
if last_error is not None:
|
||||
if attempt == max_attempts:
|
||||
break
|
||||
delay = 0.25 * (2 ** (attempt - 1))
|
||||
logger.warning(
|
||||
"Cashu transaction storage failed; retrying",
|
||||
extra={
|
||||
"type": typ,
|
||||
"request_id": request_id,
|
||||
"attempt": attempt,
|
||||
"max_attempts": max_attempts,
|
||||
"retry_delay_seconds": delay,
|
||||
},
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
|
||||
logger.critical(
|
||||
"Cashu transaction storage failed after bounded retries",
|
||||
extra={
|
||||
"type": typ,
|
||||
"request_id": request_id,
|
||||
"attempts": max_attempts,
|
||||
"error": str(last_error),
|
||||
},
|
||||
)
|
||||
if last_error is None:
|
||||
raise RuntimeError("Cashu transaction storage failed without an exception")
|
||||
raise last_error
|
||||
|
||||
|
||||
class UpstreamProviderRow(SQLModel, table=True): # type: ignore
|
||||
@@ -352,6 +440,44 @@ class RoutstrFee(SQLModel, table=True): # type: ignore
|
||||
accumulated_msats: int = Field(default=0)
|
||||
total_paid_msats: int = Field(default=0)
|
||||
last_paid_at: int | None = Field(default=None)
|
||||
payout_in_progress_msats: int = Field(default=0)
|
||||
payout_started_at: int | None = Field(default=None)
|
||||
|
||||
|
||||
class NsecState(str, Enum):
|
||||
"""Ownership state of the node's nsec — an explicit 3-state machine.
|
||||
|
||||
The single ``encrypted_nsec`` column cannot distinguish "never migrated" from
|
||||
"intentionally cleared" (both leave it empty), which let a cleared identity be
|
||||
resurrected from a stale legacy ``NSEC``. This names the three states so the
|
||||
bootstrap branches on ownership rather than inferring it:
|
||||
|
||||
* ``legacy`` — the vault has not taken ownership; a plaintext ``NSEC`` (env or
|
||||
old settings blob) may still exist and should be migrated in once.
|
||||
* ``encrypted`` — the vault owns a ciphertext; decrypt it, never re-read env.
|
||||
* ``cleared`` — the vault owns it but the operator emptied it; stay empty,
|
||||
never re-import from a stale legacy copy.
|
||||
"""
|
||||
|
||||
legacy = "legacy"
|
||||
encrypted = "encrypted"
|
||||
cleared = "cleared"
|
||||
|
||||
|
||||
class Secret(SQLModel, table=True): # type: ignore
|
||||
"""Node-level secrets, stored encrypted/hashed at rest (singleton, id=1).
|
||||
|
||||
The asymmetric column names document the encoding: ``_hash`` is one-way
|
||||
(scrypt, verify only) while ``encrypted_`` is reversible (Fernet). Per-provider
|
||||
upstream keys live on ``upstream_providers``, not here. See ``routstr.core.vault``.
|
||||
"""
|
||||
|
||||
__tablename__ = "secrets"
|
||||
id: int = Field(default=1, primary_key=True)
|
||||
admin_password_hash: str | None = Field(default=None)
|
||||
encrypted_nsec: str | None = Field(default=None)
|
||||
nsec_state: NsecState = Field(default=NsecState.legacy)
|
||||
updated_at: int | None = Field(default=None)
|
||||
|
||||
|
||||
class CliToken(SQLModel, table=True): # type: ignore
|
||||
@@ -392,18 +518,91 @@ async def get_routstr_fee(session: AsyncSession) -> RoutstrFee:
|
||||
return fee
|
||||
|
||||
|
||||
async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> None:
|
||||
async def get_secret(session: AsyncSession) -> Secret:
|
||||
secret = await session.get(Secret, 1)
|
||||
if secret is None:
|
||||
secret = Secret(id=1)
|
||||
session.add(secret)
|
||||
try:
|
||||
await session.commit()
|
||||
except IntegrityError:
|
||||
# Another worker created the singleton row between our read and
|
||||
# insert (multiple workers booting against one shared DB). Roll back
|
||||
# and read the row they committed instead of failing startup.
|
||||
await session.rollback()
|
||||
secret = await session.get(Secret, 1)
|
||||
if secret is None:
|
||||
raise
|
||||
return secret
|
||||
await session.refresh(secret)
|
||||
return secret
|
||||
|
||||
|
||||
async def set_admin_password(session: AsyncSession, password: str) -> None:
|
||||
"""Store the admin password as a one-way hash on the Secret singleton."""
|
||||
from .vault import hash_password
|
||||
|
||||
secret = await get_secret(session)
|
||||
secret.admin_password_hash = hash_password(password)
|
||||
secret.updated_at = int(time.time())
|
||||
session.add(secret)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def set_nsec(session: AsyncSession, nsec: str) -> None:
|
||||
"""Store the node's nsec, Fernet-encrypted, on the Secret singleton.
|
||||
|
||||
An empty string clears it (the node then holds no Nostr identity and signs
|
||||
no events). Either way the vault now owns the nsec, so the state moves off
|
||||
``legacy``: a cleared identity (``cleared``) must not be resurrected from a
|
||||
stale legacy ``NSEC`` on the next boot.
|
||||
"""
|
||||
from .vault import encrypt
|
||||
|
||||
secret = await get_secret(session)
|
||||
secret.encrypted_nsec = encrypt(nsec) if nsec else None
|
||||
secret.nsec_state = NsecState.encrypted if nsec else NsecState.cleared
|
||||
secret.updated_at = int(time.time())
|
||||
session.add(secret)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> bool:
|
||||
"""Checkpoint a fee payout before making the external payment."""
|
||||
stmt = (
|
||||
update(RoutstrFee)
|
||||
.where(col(RoutstrFee.id) == 1)
|
||||
.where(col(RoutstrFee.payout_in_progress_msats) == 0)
|
||||
.where(col(RoutstrFee.accumulated_msats) >= paid_msats)
|
||||
.values(
|
||||
accumulated_msats=RoutstrFee.accumulated_msats - paid_msats,
|
||||
payout_in_progress_msats=paid_msats,
|
||||
payout_started_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
||||
await session.commit()
|
||||
return result.rowcount == 1
|
||||
|
||||
|
||||
async def complete_routstr_fee_payout(
|
||||
session: AsyncSession, paid_msats: int
|
||||
) -> bool:
|
||||
"""Mark a checkpointed payout complete after the external payment succeeds."""
|
||||
stmt = (
|
||||
update(RoutstrFee)
|
||||
.where(col(RoutstrFee.id) == 1)
|
||||
.where(col(RoutstrFee.payout_in_progress_msats) == paid_msats)
|
||||
.values(
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
total_paid_msats=RoutstrFee.total_paid_msats + paid_msats,
|
||||
last_paid_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
await session.exec(stmt) # type: ignore[call-overload]
|
||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
||||
await session.commit()
|
||||
return result.rowcount == 1
|
||||
|
||||
|
||||
async def balances_for_mint_and_unit(
|
||||
|
||||
+23
-7
@@ -37,7 +37,7 @@ from .exceptions import general_exception_handler, http_exception_handler
|
||||
from .logging import get_logger, setup_logging
|
||||
from .middleware import LoggingMiddleware
|
||||
from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401
|
||||
from .settings import SettingsService
|
||||
from .settings import SettingsService, bootstrap_secrets
|
||||
from .settings import settings as global_settings
|
||||
from .version import __version__
|
||||
|
||||
@@ -85,17 +85,20 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
|
||||
# Initialize application settings (env -> computed -> DB precedence)
|
||||
async with create_session() as session:
|
||||
# Move secrets into the encrypted/hashed store and decrypt the nsec
|
||||
# into the in-memory settings BEFORE initializing settings: the
|
||||
# initialize step strips secrets from the persisted blob, so legacy
|
||||
# plaintext (env or old blob) must be migrated into the Secret store
|
||||
# first or the only copy of a blob-only secret would be lost.
|
||||
# Generates and logs an admin password on a fresh node; fails fast if
|
||||
# a stored secret can't be decrypted.
|
||||
await bootstrap_secrets(session)
|
||||
s = await SettingsService.initialize(session)
|
||||
if s.reset_reserved_balance_on_startup:
|
||||
from .db import reset_all_reserved_balances
|
||||
|
||||
await reset_all_reserved_balances(session)
|
||||
|
||||
if not s.admin_password:
|
||||
logger.warning(
|
||||
f"Admin password is not set. Visit {s.http_url or 'http://localhost:8000'}/admin to set the password."
|
||||
)
|
||||
|
||||
# Apply app metadata from settings
|
||||
try:
|
||||
app.title = s.name
|
||||
@@ -257,7 +260,20 @@ app.add_middleware(
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
expose_headers=["x-routstr-request-id", "x-cashu"],
|
||||
expose_headers=[
|
||||
"x-routstr-request-id",
|
||||
"x-cashu",
|
||||
"x-routstr-cost-msats",
|
||||
"x-routstr-cost-usd",
|
||||
"x-routstr-input-cost-msats",
|
||||
"x-routstr-output-cost-msats",
|
||||
# EHBP (Tinfoil) protocol headers must be exposed so browser clients
|
||||
# can detect and decrypt encrypted responses. Without these, the
|
||||
# browser hides them via CORS and the SDK treats the response as a
|
||||
# plaintext proxy error, returning raw ciphertext.
|
||||
"Ehbp-Response-Nonce",
|
||||
"Ehbp-Encapsulated-Key",
|
||||
],
|
||||
)
|
||||
|
||||
# Add logging middleware
|
||||
|
||||
+238
-26
@@ -3,6 +3,8 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import secrets
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
@@ -26,7 +28,6 @@ class Settings(BaseSettings):
|
||||
# Core
|
||||
upstream_base_url: str = Field(default="", env="UPSTREAM_BASE_URL")
|
||||
upstream_api_key: str = Field(default="", env="UPSTREAM_API_KEY")
|
||||
admin_password: str = Field(default="", env="ADMIN_PASSWORD")
|
||||
|
||||
# Node info
|
||||
name: str = Field(default="ARoutstrNode", env="NAME")
|
||||
@@ -132,10 +133,70 @@ def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]:
|
||||
return normalized
|
||||
|
||||
|
||||
# Secrets are credentials, not config: they live in the encrypted/hashed Secret
|
||||
# store (and decrypted in-memory for runtime use), never in the persisted
|
||||
# settings blob. ``admin_password`` is gone from the model entirely; ``nsec``
|
||||
# remains a live field but is stripped from every blob write so it is never
|
||||
# written back to plaintext. ``upstream_api_key`` is intentionally *not* here:
|
||||
# it has no encrypted home yet (it is node-scoped today but really belongs on a
|
||||
# provider), so stripping it would lose it on the next restart. It stays in the
|
||||
# blob as before; encrypting it is follow-up work. See ``bootstrap_secrets`` and
|
||||
# ``routstr.core.vault``.
|
||||
SECRET_FIELDS = frozenset({"admin_password", "nsec"})
|
||||
|
||||
|
||||
def _strip_secret_fields(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Return a copy of ``data`` without any secret fields (for persistence)."""
|
||||
return {k: v for k, v in data.items() if k not in SECRET_FIELDS}
|
||||
|
||||
|
||||
def _apply_to_live_settings(data: dict[str, Any]) -> None:
|
||||
"""Apply ``data`` onto the live ``settings`` for all in-process importers.
|
||||
|
||||
Secrets are owned exclusively by ``bootstrap_secrets``, which runs first and
|
||||
has already decrypted the authoritative nsec into memory (importing any
|
||||
legacy plaintext on the way). Never re-apply secret fields from env/blob
|
||||
here: a non-empty but stale ``NSEC`` env var would otherwise override an nsec
|
||||
the vault has taken ownership of (e.g. after the operator rotates it in the
|
||||
UI), and an empty one would wipe the live value. Skip them entirely.
|
||||
"""
|
||||
for k, v in data.items():
|
||||
if k in SECRET_FIELDS:
|
||||
continue
|
||||
setattr(settings, k, v)
|
||||
|
||||
|
||||
def _compute_primary_mint(cashu_mints: list[str]) -> str:
|
||||
return cashu_mints[0] if cashu_mints else "https://mint.minibits.cash/Bitcoin"
|
||||
|
||||
|
||||
def derive_npub_from_nsec(nsec: str) -> str | None:
|
||||
"""Derive the npub (bech32) from an nsec or 64-char hex private key, or None.
|
||||
|
||||
Parsing is delegated to :func:`routstr.nostr.listing.nsec_to_keypair`, the
|
||||
single place that knows the nsec/hex formats (and already returns ``None`` on
|
||||
any unusable input); this only bech32-encodes the resulting public key. The
|
||||
contract stays "return None on unusable input", so a bad key never crashes
|
||||
boot.
|
||||
"""
|
||||
try:
|
||||
from nostr.key import PublicKey # type: ignore
|
||||
|
||||
from ..nostr.listing import nsec_to_keypair
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
keypair = nsec_to_keypair(nsec)
|
||||
if keypair is None:
|
||||
return None
|
||||
_privkey_hex, pubkey_hex = keypair
|
||||
|
||||
try:
|
||||
return PublicKey(bytes.fromhex(pubkey_hex)).bech32()
|
||||
except (ValueError, AttributeError):
|
||||
return None
|
||||
|
||||
|
||||
def resolve_bootstrap() -> Settings:
|
||||
base = Settings() # Reads env with custom parse_env_var
|
||||
# Back-compat env mapping
|
||||
@@ -190,23 +251,9 @@ def resolve_bootstrap() -> Settings:
|
||||
pass
|
||||
# Derive NPUB from NSEC if not provided
|
||||
if not base.npub and base.nsec:
|
||||
try:
|
||||
from nostr.key import PrivateKey # type: ignore
|
||||
|
||||
if base.nsec.startswith("nsec"):
|
||||
pk = PrivateKey.from_nsec(base.nsec)
|
||||
elif len(base.nsec) == 64:
|
||||
pk = PrivateKey(bytes.fromhex(base.nsec))
|
||||
else:
|
||||
pk = None
|
||||
if pk is not None:
|
||||
try:
|
||||
base.npub = pk.public_key.bech32()
|
||||
except Exception:
|
||||
# Fallback to hex if bech32 not available
|
||||
base.npub = pk.public_key.hex()
|
||||
except Exception:
|
||||
pass
|
||||
npub = derive_npub_from_nsec(base.nsec)
|
||||
if npub:
|
||||
base.npub = npub
|
||||
if not base.cors_origins:
|
||||
base.cors_origins = ["*"]
|
||||
if not base.primary_mint:
|
||||
@@ -256,15 +303,14 @@ class SettingsService:
|
||||
text(
|
||||
"INSERT INTO settings (id, data, updated_at) VALUES (1, :data, :updated_at)"
|
||||
).bindparams(
|
||||
data=json.dumps(env_resolved.dict()),
|
||||
data=json.dumps(_strip_secret_fields(env_resolved.dict())),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
await db_session.commit()
|
||||
cls._current = settings
|
||||
# Update the existing instance in-place for all live importers
|
||||
for k, v in env_resolved.dict().items():
|
||||
setattr(settings, k, v)
|
||||
_apply_to_live_settings(env_resolved.dict())
|
||||
return cls._current
|
||||
|
||||
db_id, db_data, _updated_at = row
|
||||
@@ -291,20 +337,37 @@ class SettingsService:
|
||||
merged_dict.get("cashu_mints", [])
|
||||
)
|
||||
|
||||
if db_json_raw != merged_dict:
|
||||
# Keep npub consistent with the live nsec. bootstrap_secrets has
|
||||
# already run and holds the single authoritative nsec (decrypted from
|
||||
# the encrypted store, or freshly imported). merged_dict starts from
|
||||
# the env/blob, which may carry a STALE nsec — and therefore a stale
|
||||
# derived npub — after the vault took ownership. Derive from the live
|
||||
# value and OVERRIDE, not just fill: otherwise the node keeps the
|
||||
# vault's private key but announces the old env key's npub (npub is a
|
||||
# pure derivation of nsec, never configured independently of it).
|
||||
if settings.nsec:
|
||||
derived_npub = derive_npub_from_nsec(settings.nsec)
|
||||
if derived_npub:
|
||||
merged_dict["npub"] = derived_npub
|
||||
|
||||
# Persist without secrets; compare against the stripped target so a
|
||||
# legacy blob that still carries plaintext secrets gets rewritten
|
||||
# (and thereby sunset) even when its non-secret values are unchanged.
|
||||
persisted = _strip_secret_fields(merged_dict)
|
||||
if db_json_raw != persisted:
|
||||
await db_session.exec( # type: ignore
|
||||
text(
|
||||
"UPDATE settings SET data = :data, updated_at = :updated_at WHERE id = 1"
|
||||
).bindparams(
|
||||
data=json.dumps(merged_dict),
|
||||
data=json.dumps(persisted),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
await db_session.commit()
|
||||
|
||||
# Update the existing instance in-place for all live importers
|
||||
for k, v in merged_dict.items():
|
||||
setattr(settings, k, v)
|
||||
# (keeps the decrypted nsec live in memory).
|
||||
_apply_to_live_settings(merged_dict)
|
||||
cls._current = settings
|
||||
return cls._current
|
||||
|
||||
@@ -326,7 +389,7 @@ class SettingsService:
|
||||
text(
|
||||
"UPDATE settings SET data = :data, updated_at = :updated_at WHERE id = 1"
|
||||
).bindparams(
|
||||
data=json.dumps(candidate.dict()),
|
||||
data=json.dumps(_strip_secret_fields(candidate.dict())),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
@@ -355,3 +418,152 @@ class SettingsService:
|
||||
setattr(settings, k, v)
|
||||
cls._current = settings
|
||||
return settings
|
||||
|
||||
|
||||
async def _read_raw_settings_blob(db_session: AsyncSession) -> dict[str, Any]:
|
||||
"""Best-effort read of the raw persisted settings JSON (may not exist yet)."""
|
||||
from sqlmodel import text
|
||||
|
||||
try:
|
||||
result = await db_session.exec( # type: ignore
|
||||
text("SELECT data FROM settings WHERE id = 1")
|
||||
)
|
||||
row = result.first()
|
||||
except Exception:
|
||||
return {}
|
||||
if row is None:
|
||||
return {}
|
||||
(data_str,) = row
|
||||
try:
|
||||
data = json.loads(data_str) if isinstance(data_str, str) else dict(data_str)
|
||||
except Exception:
|
||||
return {}
|
||||
return data if isinstance(data, dict) else {}
|
||||
|
||||
|
||||
def _legacy_plaintext(
|
||||
raw_blob: dict[str, Any], env_name: str, blob_key: str
|
||||
) -> str | None:
|
||||
"""Legacy plaintext for a secret: env first, then the old settings blob."""
|
||||
env_value = os.environ.get(env_name)
|
||||
if env_value:
|
||||
return env_value
|
||||
blob_value = raw_blob.get(blob_key)
|
||||
if isinstance(blob_value, str) and blob_value:
|
||||
return blob_value
|
||||
return None
|
||||
|
||||
|
||||
async def bootstrap_secrets(db_session: AsyncSession) -> None:
|
||||
"""Move node secrets into the encrypted/hashed Secret store at startup.
|
||||
|
||||
Per secret:
|
||||
* column already set -> use it (the nsec is decrypted into the in-memory
|
||||
``settings``; a wrong ROUTSTR_SECRET_KEY surfaces as a clear fail-fast).
|
||||
* column empty but legacy plaintext exists (env, or the old settings
|
||||
blob) -> transform it (hash the password / encrypt the nsec) into the
|
||||
column.
|
||||
* nothing (admin password only) -> generate a strong random password,
|
||||
hash it, and log it once with the /admin URL.
|
||||
"""
|
||||
from cryptography.fernet import InvalidToken
|
||||
from sqlmodel import col, update
|
||||
|
||||
from . import vault
|
||||
from .db import NsecState, Secret, get_secret
|
||||
|
||||
raw_blob = await _read_raw_settings_blob(db_session)
|
||||
secret = await get_secret(db_session)
|
||||
changed = False
|
||||
|
||||
# Admin password — one-way scrypt hash.
|
||||
if secret.admin_password_hash is None:
|
||||
legacy_password = _legacy_plaintext(
|
||||
raw_blob, "ADMIN_PASSWORD", "admin_password"
|
||||
)
|
||||
if legacy_password:
|
||||
secret.admin_password_hash = vault.hash_password(legacy_password)
|
||||
changed = True
|
||||
else:
|
||||
generated = secrets.token_urlsafe(24)
|
||||
# Claim the empty slot atomically: only the worker whose UPDATE flips
|
||||
# NULL -> hash owns the generated password and announces it. On a
|
||||
# shared DB a racing worker gets rowcount 0, so it neither clobbers
|
||||
# the winner's hash (which the operator may already be using) nor
|
||||
# prints a second password that would never work.
|
||||
claim_stmt = (
|
||||
update(Secret)
|
||||
.where(col(Secret.id) == 1)
|
||||
.where(col(Secret.admin_password_hash).is_(None))
|
||||
.values(
|
||||
admin_password_hash=vault.hash_password(generated),
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
result = await db_session.exec(claim_stmt) # type: ignore[call-overload]
|
||||
await db_session.commit()
|
||||
await db_session.refresh(secret)
|
||||
if result.rowcount == 1:
|
||||
admin_url = (settings.http_url or "http://localhost:8000").rstrip("/")
|
||||
# Print to stdout rather than the logger: the operator must see
|
||||
# this once (e.g. `docker compose logs`), but it must not be
|
||||
# persisted into the on-disk log files the logger also writes to.
|
||||
print(
|
||||
"No admin password set; generated a temporary one (shown "
|
||||
f"only now): {generated}\nLog in at {admin_url}/admin and "
|
||||
"change it from the dashboard settings.",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
# Nostr nsec — reversible Fernet encryption. ``nsec_state`` is the single
|
||||
# source of truth for ownership, so "intentionally cleared" is never
|
||||
# conflated with "never migrated" (the bug the old bool could not encode).
|
||||
if secret.nsec_state == NsecState.encrypted:
|
||||
# The vault owns the identity: decrypt the ciphertext, never re-read
|
||||
# env/blob. A missing ciphertext here means the row is inconsistent (a
|
||||
# failed write or manual edit); fail fast rather than silently dropping
|
||||
# the identity and falling back to a stale legacy copy.
|
||||
if secret.encrypted_nsec is None:
|
||||
raise RuntimeError(
|
||||
"nsec_state is 'encrypted' but no ciphertext is stored; the "
|
||||
"secrets row is inconsistent. Refusing to boot rather than "
|
||||
"silently resurrecting a stale legacy NSEC."
|
||||
)
|
||||
try:
|
||||
settings.nsec = vault.decrypt(secret.encrypted_nsec)
|
||||
except InvalidToken as exc:
|
||||
raise RuntimeError(
|
||||
"Stored nsec cannot be decrypted with the current "
|
||||
"ROUTSTR_SECRET_KEY. The key changed, or this database came from "
|
||||
"another node. Restore the original ROUTSTR_SECRET_KEY to recover."
|
||||
) from exc
|
||||
elif secret.nsec_state == NsecState.cleared:
|
||||
# The operator emptied the identity via the admin API. A fresh process
|
||||
# has already reloaded a stale ``NSEC`` from env/blob into the live
|
||||
# settings (and may have derived its npub); actively clear both so the
|
||||
# cleared store wins rather than silently resurrecting the old identity.
|
||||
settings.nsec = ""
|
||||
settings.npub = ""
|
||||
else: # NsecState.legacy — the vault has not taken ownership yet.
|
||||
# Import any legacy plaintext (env, or the old settings blob) exactly
|
||||
# once. Encryption at rest is mandatory, but a missing key is
|
||||
# provisioned, not fatal: vault.encrypt generates and persists a master
|
||||
# key (with a loud one-time operator notice) when none was supplied, so
|
||||
# an upgrading node keeps running. The nsec is never stored in plaintext.
|
||||
legacy_nsec = _legacy_plaintext(raw_blob, "NSEC", "nsec")
|
||||
if legacy_nsec:
|
||||
secret.encrypted_nsec = vault.encrypt(legacy_nsec)
|
||||
secret.nsec_state = NsecState.encrypted
|
||||
settings.nsec = legacy_nsec
|
||||
changed = True
|
||||
|
||||
# Derive npub from whatever nsec we now hold, if not already known.
|
||||
if settings.nsec and not settings.npub:
|
||||
npub = derive_npub_from_nsec(settings.nsec)
|
||||
if npub:
|
||||
settings.npub = npub
|
||||
|
||||
if changed:
|
||||
secret.updated_at = int(time.time())
|
||||
db_session.add(secret)
|
||||
await db_session.commit()
|
||||
|
||||
@@ -0,0 +1,320 @@
|
||||
"""Encrypt/hash/fingerprint helpers for secrets at rest (issue #553).
|
||||
|
||||
Thin wrapper over ``cryptography`` so nothing else in the codebase touches
|
||||
Fernet/scrypt/HMAC directly:
|
||||
|
||||
- :func:`encrypt`/:func:`decrypt` — Fernet symmetric encryption, keyed by the
|
||||
mandatory master key. Ciphertext is self-describing (``fernet:v1:`` prefix) so
|
||||
a value can be told apart from legacy plaintext and so reading it under the
|
||||
wrong key surfaces as a hard error rather than silent corruption.
|
||||
- :func:`hash_password`/:func:`verify_password` — salted scrypt hashing. This is
|
||||
*key-independent*: it never reads the master key, so password login and the
|
||||
recovery script keep working even when the key is missing.
|
||||
|
||||
Key custody is flexible but encryption is not optional. The key comes from the
|
||||
``ROUTSTR_SECRET_KEY`` env var, else a persisted key file
|
||||
(``ROUTSTR_SECRET_KEY_FILE``, defaulting beside the SQLite database so it persists
|
||||
on the same volume as the data); when neither is set, :func:`encrypt` generates
|
||||
one to the key file and prints a one-time notice, so an existing node upgrades
|
||||
without breaking instead of refusing to boot. Reading is strict — :func:`decrypt`
|
||||
never generates a key (a new key could not match existing ciphertext) and fails
|
||||
fast with the generation command when none is configured. A malformed
|
||||
``ROUTSTR_SECRET_KEY`` is an operator error and always fails fast.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import os
|
||||
import secrets
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from cryptography.fernet import Fernet
|
||||
from sqlalchemy.engine import make_url
|
||||
from sqlalchemy.exc import ArgumentError
|
||||
|
||||
_PREFIX = "fernet:v1:"
|
||||
_GEN_COMMAND = (
|
||||
'python -c "from cryptography.fernet import Fernet; '
|
||||
'print(Fernet.generate_key().decode())"'
|
||||
)
|
||||
|
||||
# Where an auto-generated master key is persisted when the operator supplies no
|
||||
# ``ROUTSTR_SECRET_KEY``. Defaults beside the SQLite database so it rides whatever
|
||||
# volume already persists the data (a container recreate would otherwise generate
|
||||
# a fresh key and be unable to decrypt existing secrets); falls back to the
|
||||
# working directory when the DB location is unknown. Override the exact path with
|
||||
# ``ROUTSTR_SECRET_KEY_FILE``.
|
||||
_KEY_FILE_ENV = "ROUTSTR_SECRET_KEY_FILE"
|
||||
_DEFAULT_KEY_FILE = "routstr_secret.key"
|
||||
|
||||
# Minimum admin-password length, enforced wherever a password is set/changed
|
||||
# (admin endpoints + the recovery script) so the policy lives in one place.
|
||||
MIN_PASSWORD_LENGTH = 8
|
||||
|
||||
# scrypt parameters; packed into each hash so verification is parameter-free.
|
||||
_SCRYPT_N = 2**14
|
||||
_SCRYPT_R = 8
|
||||
_SCRYPT_P = 1
|
||||
_SCRYPT_DKLEN = 32
|
||||
_SCRYPT_SALT_BYTES = 16
|
||||
|
||||
|
||||
def _database_dir() -> Path | None:
|
||||
"""Directory of the SQLite database file, or ``None`` when it has no on-disk
|
||||
location (a non-SQLite URL or ``:memory:``).
|
||||
|
||||
Read from ``DATABASE_URL`` at call time and parsed here rather than importing
|
||||
``routstr.core.db`` — that module builds the engine at import, which the
|
||||
crypto layer must not drag in. Mirrors db.py's ``DATABASE_URL`` default.
|
||||
"""
|
||||
url_str = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db")
|
||||
try:
|
||||
url = make_url(url_str)
|
||||
except ArgumentError:
|
||||
return None
|
||||
if url.get_backend_name() != "sqlite" or not url.database:
|
||||
return None
|
||||
if url.database == ":memory:":
|
||||
return None
|
||||
return Path(url.database).parent
|
||||
|
||||
|
||||
def _key_file_path() -> Path:
|
||||
"""Where the auto-generated master key is read from / written to.
|
||||
|
||||
``ROUTSTR_SECRET_KEY_FILE`` wins; otherwise the key sits beside the SQLite
|
||||
database so it persists on the same volume as the data, falling back to the
|
||||
working directory when the DB location is unknown.
|
||||
"""
|
||||
override = os.environ.get(_KEY_FILE_ENV)
|
||||
if override:
|
||||
return Path(override)
|
||||
directory = _database_dir()
|
||||
return (directory or Path()) / _DEFAULT_KEY_FILE
|
||||
|
||||
|
||||
def _read_key_file(path: Path) -> str | None:
|
||||
try:
|
||||
stored = path.read_text().strip()
|
||||
except OSError:
|
||||
return None
|
||||
if not stored:
|
||||
return None
|
||||
_repair_key_file_perms(path)
|
||||
return stored
|
||||
|
||||
|
||||
def _repair_key_file_perms(path: Path) -> None:
|
||||
# A master key must never be group/other-readable. On POSIX, tighten loose
|
||||
# permissions to owner-only (0600) rather than trust — or hard-fail on — a
|
||||
# world-readable key; a friendlier repair keeps an upgrading node booting.
|
||||
if os.name != "posix":
|
||||
return
|
||||
try:
|
||||
mode = path.stat().st_mode
|
||||
except OSError:
|
||||
return
|
||||
if mode & 0o077:
|
||||
try:
|
||||
os.chmod(path, 0o600)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def _load_secret_key() -> str | None:
|
||||
"""The configured key without provisioning: env var, then the key file."""
|
||||
return os.environ.get("ROUTSTR_SECRET_KEY") or _read_key_file(_key_file_path())
|
||||
|
||||
|
||||
def _warn_generated_key(path: Path) -> None:
|
||||
# stdout, not the logger: the operator must see this once (e.g. in
|
||||
# ``docker compose logs``), but it must never be persisted into the on-disk
|
||||
# log files the logger also writes. Mirrors the generated-admin-password
|
||||
# notice so an upgrade cannot silently create an unbacked key.
|
||||
print(
|
||||
"No ROUTSTR_SECRET_KEY was set; generated one to encrypt node secrets at "
|
||||
f"rest and saved it to {path}.\n"
|
||||
"!! BACK UP THIS FILE. If it is lost, the encrypted secrets cannot be "
|
||||
"recovered and will have to be re-entered.\n"
|
||||
"To manage the key yourself (e.g. from a secrets manager) set "
|
||||
"ROUTSTR_SECRET_KEY in the environment instead; the value is in the file "
|
||||
"above.",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def _generate_and_persist_key(path: Path) -> str:
|
||||
"""Generate a Fernet key, persist it owner-only and atomically, warn once.
|
||||
|
||||
The key is written to a temp file in the same directory, flushed durably,
|
||||
then ``os.link``-ed into place. ``os.link`` publishes the complete file in a
|
||||
single atomic step — a crash mid-write leaves only the temp file (which is
|
||||
removed), never a half-written or empty key at the final path that a later
|
||||
boot would read as corrupt. It also refuses to overwrite an existing key, so
|
||||
a racing worker that generated first keeps ownership (secrets may already be
|
||||
encrypted under its key); the loser adopts that key instead of clobbering it.
|
||||
"""
|
||||
key = Fernet.generate_key().decode()
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
fd, tmp_name = tempfile.mkstemp(
|
||||
dir=path.parent, prefix=".routstr_secret.", suffix=".tmp"
|
||||
)
|
||||
tmp = Path(tmp_name)
|
||||
try:
|
||||
with os.fdopen(fd, "w") as handle:
|
||||
handle.write(key) # mkstemp already created it 0600
|
||||
handle.flush()
|
||||
os.fsync(handle.fileno())
|
||||
try:
|
||||
os.link(tmp, path)
|
||||
except FileExistsError:
|
||||
# A concurrent worker linked its key in first; adopt theirs rather
|
||||
# than clobber a key that secrets may already be encrypted under.
|
||||
existing = _read_key_file(path)
|
||||
if existing:
|
||||
return existing
|
||||
raise
|
||||
_fsync_dir(path.parent)
|
||||
finally:
|
||||
tmp.unlink(missing_ok=True)
|
||||
_warn_generated_key(path)
|
||||
return key
|
||||
|
||||
|
||||
def _fsync_dir(directory: Path) -> None:
|
||||
# Persist the new directory entry so the linked key survives a crash right
|
||||
# after publish. Best-effort: not every platform lets you fsync a directory.
|
||||
try:
|
||||
dir_fd = os.open(directory, os.O_RDONLY)
|
||||
except OSError:
|
||||
return
|
||||
try:
|
||||
os.fsync(dir_fd)
|
||||
except OSError:
|
||||
pass
|
||||
finally:
|
||||
os.close(dir_fd)
|
||||
|
||||
|
||||
def ensure_secret_key() -> str:
|
||||
"""Return the master key, provisioning one if the operator supplied none.
|
||||
|
||||
Precedence: the ``ROUTSTR_SECRET_KEY`` env var, then the persisted key file,
|
||||
otherwise a freshly generated key written to the key file (with a one-time
|
||||
operator notice). This keeps encryption at rest mandatory while letting an
|
||||
existing node upgrade without setting a key first. A malformed env key is
|
||||
left to fail at :func:`get_fernet` — it is an operator error, not an unset
|
||||
key, so it must not trigger silent self-provisioning.
|
||||
"""
|
||||
env_key = os.environ.get("ROUTSTR_SECRET_KEY")
|
||||
if env_key:
|
||||
return env_key
|
||||
path = _key_file_path()
|
||||
return _read_key_file(path) or _generate_and_persist_key(path)
|
||||
|
||||
|
||||
def _fernet_from_key(key: str) -> Fernet:
|
||||
try:
|
||||
return Fernet(key.encode())
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise RuntimeError(
|
||||
"ROUTSTR_SECRET_KEY is malformed; it must be a url-safe base64 "
|
||||
"32-byte Fernet key. Generate one with:\n " + _GEN_COMMAND
|
||||
) from exc
|
||||
|
||||
|
||||
def get_fernet() -> Fernet:
|
||||
"""Build a :class:`Fernet` from the configured key (env var or key file).
|
||||
|
||||
Strict: this never generates a key, so already-encrypted ciphertext is never
|
||||
shadowed by a fresh key. A read with no key configured fails fast with the
|
||||
generation command.
|
||||
"""
|
||||
key = _load_secret_key()
|
||||
if not key:
|
||||
raise RuntimeError(
|
||||
"ROUTSTR_SECRET_KEY is not set. It is required to encrypt secrets at "
|
||||
"rest. Generate one with:\n " + _GEN_COMMAND
|
||||
)
|
||||
return _fernet_from_key(key)
|
||||
|
||||
|
||||
def encrypt(plaintext: str) -> str:
|
||||
"""Encrypt ``plaintext`` into a self-describing ``fernet:v1:`` token.
|
||||
|
||||
Provisions a master key (env var, key file, or a freshly generated one) so an
|
||||
upgrading node never has to set one before its first secret is stored; the
|
||||
value is always encrypted, never persisted in plaintext.
|
||||
"""
|
||||
fernet = _fernet_from_key(ensure_secret_key())
|
||||
return _PREFIX + fernet.encrypt(plaintext.encode()).decode()
|
||||
|
||||
|
||||
def is_encrypted(value: str) -> bool:
|
||||
"""True if ``value`` carries the ``fernet:v1:`` prefix this module emits."""
|
||||
return value.startswith(_PREFIX)
|
||||
|
||||
|
||||
def decrypt(ciphertext: str) -> str:
|
||||
"""Decrypt a ``fernet:v1:`` token.
|
||||
|
||||
Raises ``ValueError`` for an unprefixed value (so legacy plaintext is never
|
||||
mistaken for ciphertext) and ``InvalidToken`` when the value was written
|
||||
under a different ``ROUTSTR_SECRET_KEY``.
|
||||
"""
|
||||
if not is_encrypted(ciphertext):
|
||||
raise ValueError("value is not fernet:v1: ciphertext")
|
||||
token = ciphertext[len(_PREFIX) :]
|
||||
return get_fernet().decrypt(token.encode()).decode()
|
||||
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
"""Salted scrypt hash, self-describing as ``scrypt:n:r:p:salt:hash``."""
|
||||
salt = secrets.token_bytes(_SCRYPT_SALT_BYTES)
|
||||
derived = hashlib.scrypt(
|
||||
password.encode(),
|
||||
salt=salt,
|
||||
n=_SCRYPT_N,
|
||||
r=_SCRYPT_R,
|
||||
p=_SCRYPT_P,
|
||||
dklen=_SCRYPT_DKLEN,
|
||||
)
|
||||
return ":".join(
|
||||
[
|
||||
"scrypt",
|
||||
str(_SCRYPT_N),
|
||||
str(_SCRYPT_R),
|
||||
str(_SCRYPT_P),
|
||||
base64.b64encode(salt).decode(),
|
||||
base64.b64encode(derived).decode(),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def verify_password(password: str, stored: str) -> bool:
|
||||
"""Constant-time check of ``password`` against a :func:`hash_password` value."""
|
||||
try:
|
||||
scheme, n, r, p, salt_b64, hash_b64 = stored.split(":")
|
||||
if scheme != "scrypt":
|
||||
return False
|
||||
n_int, r_int, p_int = int(n), int(r), int(p)
|
||||
# Cap the work factor at the parameters this module emits. scrypt's
|
||||
# memory cost grows with N*r, so an oversized N/r in a tampered or
|
||||
# corrupt stored hash could turn a single login into an OOM/DoS.
|
||||
if n_int > _SCRYPT_N or r_int > _SCRYPT_R or p_int > _SCRYPT_P:
|
||||
return False
|
||||
salt = base64.b64decode(salt_b64)
|
||||
expected = base64.b64decode(hash_b64)
|
||||
derived = hashlib.scrypt(
|
||||
password.encode(),
|
||||
salt=salt,
|
||||
n=n_int,
|
||||
r=r_int,
|
||||
p=p_int,
|
||||
dklen=len(expected),
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
return False
|
||||
return hmac.compare_digest(derived, expected)
|
||||
@@ -1,4 +1,5 @@
|
||||
import math
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from pydantic.v1 import BaseModel
|
||||
|
||||
@@ -7,6 +8,9 @@ from ..core.settings import settings
|
||||
from .price import sats_usd_price
|
||||
from .usage import normalize_usage, parse_token_count
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .models import Model
|
||||
|
||||
__all__ = [
|
||||
"CostData",
|
||||
"CostDataError",
|
||||
@@ -66,12 +70,23 @@ def _empty_cost(cls: type[CostData] = CostData) -> CostData:
|
||||
async def calculate_cost(
|
||||
response_data: dict,
|
||||
max_cost: int,
|
||||
model_obj: "Model | None" = None,
|
||||
provider_fee: float | None = None,
|
||||
) -> CostData | MaxCostData | CostDataError:
|
||||
"""Calculate the cost of an API request based on token usage.
|
||||
|
||||
Args:
|
||||
response_data: Response data containing usage information
|
||||
max_cost: Maximum cost in millisats
|
||||
model_obj: The model that actually served the request. When given,
|
||||
its pricing is billed directly; without it, pricing is re-derived
|
||||
from the response's model string via the alias map, which resolves
|
||||
to the best-ranked candidate — not necessarily the serving one.
|
||||
provider_fee: The serving provider's fee multiplier, applied on the
|
||||
USD-cost path and the litellm pricing fallback (configured model
|
||||
pricing already carries the fee baked in). Without it, the fee is
|
||||
re-derived from the response's model string, which yields the
|
||||
best-ranked provider's fee.
|
||||
|
||||
Returns:
|
||||
Cost data or error information
|
||||
@@ -157,11 +172,16 @@ async def calculate_cost(
|
||||
},
|
||||
)
|
||||
try:
|
||||
cost_details = usage_data.get("cost_details", {})
|
||||
if not isinstance(cost_details, dict):
|
||||
cost_details = {}
|
||||
input_usd = _coerce_usd(
|
||||
usage_data.get("cost_details", {}).get("input_cost", 0)
|
||||
cost_details.get("input_cost")
|
||||
or cost_details.get("upstream_inference_prompt_cost")
|
||||
)
|
||||
output_usd = _coerce_usd(
|
||||
usage_data.get("cost_details", {}).get("output_cost", 0)
|
||||
cost_details.get("output_cost")
|
||||
or cost_details.get("upstream_inference_completions_cost")
|
||||
)
|
||||
return _calculate_from_usd_cost(
|
||||
usd_cost,
|
||||
@@ -172,6 +192,7 @@ async def calculate_cost(
|
||||
cache_creation_tokens,
|
||||
output_tokens,
|
||||
response_data,
|
||||
provider_fee,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
@@ -185,7 +206,7 @@ async def calculate_cost(
|
||||
|
||||
# Fall back to token-based pricing
|
||||
try:
|
||||
pricing_rates = _get_pricing_rates(response_data)
|
||||
pricing_rates = _get_pricing_rates(response_data, model_obj, provider_fee)
|
||||
except ValueError as e:
|
||||
return CostDataError(message=str(e), code="pricing_error")
|
||||
|
||||
@@ -259,7 +280,17 @@ def _coerce_usd(value: object) -> float:
|
||||
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
||||
"""Resolve USD cost with clear priority order.
|
||||
|
||||
Priority: cost_details.total_cost → total_cost → cost (in both usage and response).
|
||||
Priority:
|
||||
|
||||
1. ``cost_details.total_cost``
|
||||
2. ``cost_details.upstream_inference_cost`` (BYOK — see below)
|
||||
3. ``total_cost`` → ``cost`` (in both usage and response)
|
||||
|
||||
**BYOK path (PPQ.AI):** when ``is_byok`` is true the ``usage.cost`` field
|
||||
is only a small (~5 %) routing fee, not the inference cost. The real cost
|
||||
lives in ``cost_details.upstream_inference_cost`` and the provider's
|
||||
balance is debited by ``upstream_inference_cost + byok_fee``. Billing just
|
||||
the fee under-charges by ~20×.
|
||||
"""
|
||||
cost_details = usage_data.get("cost_details")
|
||||
if isinstance(cost_details, dict):
|
||||
@@ -267,6 +298,18 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
||||
if cost > 0:
|
||||
return cost
|
||||
|
||||
# PPQ.AI BYOK: upstream_inference_cost is the real inference cost;
|
||||
# usage.cost is only a ~5 % BYOK routing fee. Bill the sum — what PPQ
|
||||
# actually deducts from the balance. For non-BYOK providers (e.g.
|
||||
# OpenRouter) usage.cost already equals upstream_inference_cost, so we
|
||||
# fall through to the normal ``cost`` lookup below.
|
||||
upstream_cost = _coerce_usd(
|
||||
cost_details.get("upstream_inference_cost")
|
||||
)
|
||||
if upstream_cost > 0 and usage_data.get("is_byok"):
|
||||
byok_fee = _coerce_usd(usage_data.get("cost"))
|
||||
return upstream_cost + byok_fee
|
||||
|
||||
for source in [usage_data, response_data]:
|
||||
if not isinstance(source, dict):
|
||||
continue
|
||||
@@ -280,55 +323,106 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
||||
|
||||
def _get_pricing_rates(
|
||||
response_data: dict,
|
||||
model_obj: "Model | None",
|
||||
provider_fee: float | None,
|
||||
) -> tuple[float, float, float, float] | None:
|
||||
"""Get model-based pricing rates or None if using fixed pricing.
|
||||
"""Get configured rates, falling back to LiteLLM's model cost map.
|
||||
|
||||
Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate)
|
||||
The served ``model_obj`` (when the caller has it) is billed directly;
|
||||
otherwise the response's model string is resolved through the alias map,
|
||||
which yields the best-ranked candidate rather than the serving one.
|
||||
|
||||
Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate).
|
||||
``None`` means configured fixed pricing should be used by the caller.
|
||||
"""
|
||||
if settings.fixed_pricing:
|
||||
if settings.fixed_pricing and (
|
||||
settings.fixed_per_1k_input_tokens
|
||||
or settings.fixed_per_1k_output_tokens
|
||||
):
|
||||
return None
|
||||
|
||||
from ..proxy import get_model_instance
|
||||
from .models import litellm_cost_entry
|
||||
|
||||
response_model = response_data.get("model", "")
|
||||
model_obj = get_model_instance(response_model)
|
||||
|
||||
if not model_obj:
|
||||
logger.error("Invalid model in response", extra={"response_model": response_model})
|
||||
raise ValueError(f"Invalid model: {response_model}")
|
||||
|
||||
if not model_obj.sats_pricing:
|
||||
logger.error(
|
||||
"Model pricing not defined",
|
||||
extra={"model": response_model, "model_id": response_model},
|
||||
if model_obj is None:
|
||||
logger.warning(
|
||||
"Settling without routed model identity — re-deriving pricing "
|
||||
"from the response's model string via the alias map",
|
||||
extra={"response_model": response_model},
|
||||
)
|
||||
raise ValueError("Model pricing not defined")
|
||||
model_obj = get_model_instance(response_model)
|
||||
|
||||
try:
|
||||
mspp = float(model_obj.sats_pricing.prompt)
|
||||
mspc = float(model_obj.sats_pricing.completion)
|
||||
mscr = float(model_obj.sats_pricing.input_cache_read or 0)
|
||||
mscw = float(model_obj.sats_pricing.input_cache_write or 0)
|
||||
if model_obj and model_obj.sats_pricing:
|
||||
try:
|
||||
mspp = float(model_obj.sats_pricing.prompt)
|
||||
mspc = float(model_obj.sats_pricing.completion)
|
||||
mscr = float(model_obj.sats_pricing.input_cache_read or 0)
|
||||
mscw = float(model_obj.sats_pricing.input_cache_write or 0)
|
||||
|
||||
mspp_1k = mspp * 1_000_000.0
|
||||
mspc_1k = mspc * 1_000_000.0
|
||||
mscr_1k = mscr * 1_000_000.0 if mscr > 0 else mspp_1k
|
||||
mscw_1k = mscw * 1_000_000.0 if mscw > 0 else mspp_1k
|
||||
mspp_1k = mspp * 1_000_000.0
|
||||
mspc_1k = mspc * 1_000_000.0
|
||||
mscr_1k = mscr * 1_000_000.0 if mscr > 0 else mspp_1k
|
||||
mscw_1k = mscw * 1_000_000.0 if mscw > 0 else mspp_1k
|
||||
source = "configured"
|
||||
except Exception as e:
|
||||
logger.error("Invalid pricing data", extra={"error": str(e)})
|
||||
raise ValueError("Invalid pricing data") from e
|
||||
else:
|
||||
pricing_model = (
|
||||
model_obj.forwarded_model_id if model_obj else None
|
||||
) or response_model
|
||||
pricing = litellm_cost_entry(pricing_model)
|
||||
if pricing is None:
|
||||
logger.error(
|
||||
"Model pricing not found in configured models or LiteLLM",
|
||||
extra={
|
||||
"response_model": response_model,
|
||||
"pricing_model": pricing_model,
|
||||
},
|
||||
)
|
||||
raise ValueError(f"Pricing not found for model: {response_model}")
|
||||
|
||||
logger.info(
|
||||
"Applied model-specific pricing",
|
||||
extra={
|
||||
"model": response_model,
|
||||
"input_price_msats_per_1k": mspp_1k,
|
||||
"output_price_msats_per_1k": mspc_1k,
|
||||
"cache_read_price_msats_per_1k": mscr_1k,
|
||||
"cache_write_price_msats_per_1k": mscw_1k,
|
||||
},
|
||||
input_usd = _coerce_usd(pricing.get("input_cost_per_token"))
|
||||
output_usd = _coerce_usd(pricing.get("output_cost_per_token"))
|
||||
if input_usd <= 0 or output_usd <= 0:
|
||||
raise ValueError(f"Incomplete LiteLLM pricing for model: {pricing_model}")
|
||||
|
||||
if provider_fee is None:
|
||||
provider_fee = _resolve_provider_fee(response_model)
|
||||
usd_per_sat = sats_usd_price()
|
||||
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||
mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||
cache_read_usd = _coerce_usd(
|
||||
pricing.get("cache_read_input_token_cost")
|
||||
)
|
||||
return mspp_1k, mspc_1k, mscr_1k, mscw_1k
|
||||
except Exception as e:
|
||||
logger.error("Invalid pricing data", extra={"error": str(e)})
|
||||
raise ValueError("Invalid pricing data") from e
|
||||
cache_write_usd = _coerce_usd(
|
||||
pricing.get("cache_creation_input_token_cost")
|
||||
)
|
||||
mscr_1k = (
|
||||
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||
if cache_read_usd > 0
|
||||
else mspp_1k
|
||||
)
|
||||
mscw_1k = (
|
||||
cache_write_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||
if cache_write_usd > 0
|
||||
else mspp_1k
|
||||
)
|
||||
source = "litellm"
|
||||
|
||||
logger.info(
|
||||
"Applied model-specific pricing",
|
||||
extra={
|
||||
"model": response_model,
|
||||
"pricing_source": source,
|
||||
"input_price_msats_per_1k": mspp_1k,
|
||||
"output_price_msats_per_1k": mspc_1k,
|
||||
"cache_read_price_msats_per_1k": mscr_1k,
|
||||
"cache_write_price_msats_per_1k": mscw_1k,
|
||||
},
|
||||
)
|
||||
return mspp_1k, mspc_1k, mscr_1k, mscw_1k
|
||||
|
||||
|
||||
def _resolve_provider_fee(model_id: str) -> float:
|
||||
@@ -356,9 +450,11 @@ def _calculate_from_usd_cost(
|
||||
cache_creation_tokens: int,
|
||||
output_tokens: int,
|
||||
response_data: dict,
|
||||
provider_fee: float | None,
|
||||
) -> CostData:
|
||||
"""Calculate cost from USD figures, deriving input/output split from tokens."""
|
||||
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
|
||||
if provider_fee is None:
|
||||
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
|
||||
usd_cost = usd_cost * provider_fee
|
||||
input_usd = input_usd * provider_fee
|
||||
output_usd = output_usd * provider_fee
|
||||
@@ -367,8 +463,12 @@ def _calculate_from_usd_cost(
|
||||
cost_in_msats = math.ceil(cost_in_sats * 1000)
|
||||
|
||||
if input_usd > 0 or output_usd > 0:
|
||||
input_msats = int((input_usd * sats_per_usd) * 1000)
|
||||
output_msats = int((output_usd * sats_per_usd) * 1000)
|
||||
# The total is the authoritative billed amount. Allocating that integer
|
||||
# total proportionally avoids losing sub-millisatoshi remainders when
|
||||
# input and output components are each truncated independently.
|
||||
component_usd = input_usd + output_usd
|
||||
input_msats = math.floor(cost_in_msats * input_usd / component_usd)
|
||||
output_msats = cost_in_msats - input_msats
|
||||
else:
|
||||
effective_input_tokens = (
|
||||
input_tokens + cache_read_tokens + cache_creation_tokens
|
||||
@@ -381,12 +481,37 @@ def _calculate_from_usd_cost(
|
||||
)
|
||||
output_msats = cost_in_msats - input_msats
|
||||
|
||||
# Allocate cache read/creation msats proportionally from the input portion.
|
||||
# Without this, cache_read_msats and cache_creation_msats are always 0 in
|
||||
# the USD-cost path, and the SDK cannot report cache cost breakdowns.
|
||||
cache_read_msats = 0
|
||||
cache_creation_msats = 0
|
||||
if cache_read_tokens > 0 or cache_creation_tokens > 0:
|
||||
cache_tokens = cache_read_tokens + cache_creation_tokens
|
||||
regular_input_tokens = input_tokens
|
||||
total_input_tokens = regular_input_tokens + cache_tokens
|
||||
if total_input_tokens > 0:
|
||||
# Split input_msats into regular input, cache_read, cache_creation
|
||||
cache_read_msats = (
|
||||
int(input_msats * cache_read_tokens / total_input_tokens)
|
||||
if cache_read_tokens > 0
|
||||
else 0
|
||||
)
|
||||
cache_creation_msats = (
|
||||
int(input_msats * cache_creation_tokens / total_input_tokens)
|
||||
if cache_creation_tokens > 0
|
||||
else 0
|
||||
)
|
||||
input_msats = input_msats - cache_read_msats - cache_creation_msats
|
||||
|
||||
logger.info(
|
||||
"Using cost from usage data/details",
|
||||
extra={
|
||||
"usd_cost": usd_cost,
|
||||
"cost_in_sats": cost_in_sats,
|
||||
"cost_in_msats": cost_in_msats,
|
||||
"cache_read_msats": cache_read_msats,
|
||||
"cache_creation_msats": cache_creation_msats,
|
||||
"model": response_data.get("model", "unknown"),
|
||||
},
|
||||
)
|
||||
@@ -401,8 +526,8 @@ def _calculate_from_usd_cost(
|
||||
output_tokens=output_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_msats=0,
|
||||
cache_creation_msats=0,
|
||||
cache_read_msats=cache_read_msats,
|
||||
cache_creation_msats=cache_creation_msats,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,11 +1,17 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
from typing import TypedDict
|
||||
|
||||
import httpx
|
||||
from cashu.wallet.wallet import Proof, Wallet
|
||||
|
||||
# The Cashu library issues POST /v1/melt/bolt11 with timeout=None, so a hung or
|
||||
# very slow mint can block a melt (and any caller, e.g. the payout loop)
|
||||
# indefinitely. Bound it here so callers fail instead of hanging forever.
|
||||
MELT_TIMEOUT_SECONDS = 60
|
||||
|
||||
try:
|
||||
from bech32 import bech32_decode, convertbits # type: ignore
|
||||
except ModuleNotFoundError: # pragma: no cover – allow runtime miss
|
||||
@@ -220,10 +226,18 @@ async def raw_send_to_lnurl(
|
||||
if amount:
|
||||
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
|
||||
|
||||
_ = await wallet.melt(
|
||||
proofs=proofs,
|
||||
invoice=bolt11_invoice,
|
||||
fee_reserve_sat=melt_quote_resp.fee_reserve,
|
||||
quote_id=melt_quote_resp.quote,
|
||||
)
|
||||
try:
|
||||
_ = await asyncio.wait_for(
|
||||
wallet.melt(
|
||||
proofs=proofs,
|
||||
invoice=bolt11_invoice,
|
||||
fee_reserve_sat=melt_quote_resp.fee_reserve,
|
||||
quote_id=melt_quote_resp.quote,
|
||||
),
|
||||
timeout=MELT_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise LNURLError(
|
||||
f"Melt timed out after {MELT_TIMEOUT_SECONDS}s (mint unresponsive)"
|
||||
) from e
|
||||
return final_amount
|
||||
|
||||
+203
-64
@@ -29,6 +29,7 @@ from .payment.helpers import (
|
||||
)
|
||||
from .payment.models import Model
|
||||
from .upstream import BaseUpstreamProvider
|
||||
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
|
||||
from .upstream.helpers import init_upstreams
|
||||
from .upstream.request_correction import correct_request, extract_error_message
|
||||
|
||||
@@ -36,10 +37,9 @@ logger = get_logger(__name__)
|
||||
proxy_router = APIRouter()
|
||||
|
||||
_upstreams: list[BaseUpstreamProvider] = []
|
||||
_model_instances: dict[str, Model] = {} # All aliases -> Model
|
||||
_provider_map: dict[
|
||||
str, list[BaseUpstreamProvider]
|
||||
] = {} # All aliases -> List[Provider]
|
||||
str, list[tuple[Model, BaseUpstreamProvider]]
|
||||
] = {} # All aliases -> sorted [(candidate Model, its Provider)]
|
||||
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
|
||||
|
||||
|
||||
@@ -71,32 +71,44 @@ def get_upstreams() -> list[BaseUpstreamProvider]:
|
||||
return _upstreams
|
||||
|
||||
|
||||
def get_model_instance(model_id: str) -> Model | None:
|
||||
"""Get Model instance by ID from global cache."""
|
||||
def get_candidates(
|
||||
model_id: str,
|
||||
) -> list[tuple[Model, BaseUpstreamProvider]] | None:
|
||||
"""Get the sorted (model, provider) candidate list for a model ID.
|
||||
|
||||
Each provider is paired with its own model for the alias, so routing can
|
||||
forward and bill the candidate that actually serves. Version suffixes
|
||||
(e.g. ``-20251222``) are stripped as a retry when the exact ID is
|
||||
unknown, since upstreams may return a specific version of a base model
|
||||
we track.
|
||||
"""
|
||||
if not model_id:
|
||||
return None
|
||||
|
||||
model_id_lower = model_id.lower()
|
||||
# Try exact match first
|
||||
if model := _model_instances.get(model_id_lower):
|
||||
return model
|
||||
if candidates := _provider_map.get(model_id_lower):
|
||||
return candidates
|
||||
|
||||
# Try stripping common version suffixes (e.g., -20251222)
|
||||
# This handles cases where upstream returns a specific version
|
||||
# but we only track the base model name.
|
||||
import re
|
||||
|
||||
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
|
||||
if base_model_id != model_id_lower:
|
||||
if model := _model_instances.get(base_model_id):
|
||||
return model
|
||||
if candidates := _provider_map.get(base_model_id):
|
||||
return candidates
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def get_model_instance(model_id: str) -> Model | None:
|
||||
"""Get the best-ranked Model instance for a model ID."""
|
||||
candidates = get_candidates(model_id)
|
||||
return candidates[0][0] if candidates else None
|
||||
|
||||
|
||||
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
|
||||
"""Get UpstreamProvider list for model ID from global cache."""
|
||||
return _provider_map.get(model_id.lower())
|
||||
"""Get the sorted UpstreamProvider list for a model ID."""
|
||||
candidates = get_candidates(model_id)
|
||||
return [provider for _, provider in candidates] if candidates else None
|
||||
|
||||
|
||||
def get_unique_models() -> list[Model]:
|
||||
@@ -104,11 +116,39 @@ def get_unique_models() -> list[Model]:
|
||||
return list(_unique_models.values())
|
||||
|
||||
|
||||
def _is_tinfoil_attestation_path(path: str) -> bool:
|
||||
"""Return True for exact Tinfoil attestation routes, with optional slash."""
|
||||
return path in {
|
||||
"attestation",
|
||||
"attestation/",
|
||||
"tee/attestation",
|
||||
"tee/attestation/",
|
||||
}
|
||||
|
||||
|
||||
def _select_unauthenticated_get_upstreams(
|
||||
path: str, upstreams: list[BaseUpstreamProvider]
|
||||
) -> list[BaseUpstreamProvider]:
|
||||
"""Select upstream candidates for unauthenticated GET bypass paths.
|
||||
|
||||
Tinfoil attestation endpoints are provider-specific. Trying every enabled
|
||||
upstream can return an unrelated provider's 404 before Tinfoil is reached,
|
||||
so route those paths only to Tinfoil providers.
|
||||
"""
|
||||
if _is_tinfoil_attestation_path(path):
|
||||
return [
|
||||
upstream
|
||||
for upstream in upstreams
|
||||
if getattr(upstream, "provider_type", None) == "tinfoil"
|
||||
]
|
||||
return upstreams
|
||||
|
||||
|
||||
async def refresh_model_maps() -> None:
|
||||
"""Refresh global model and provider maps using the cost-based algorithm."""
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
global _model_instances, _provider_map, _unique_models
|
||||
global _provider_map, _unique_models
|
||||
|
||||
async with create_session() as session:
|
||||
# Fetch all providers with their models in a single logical operation
|
||||
@@ -118,22 +158,23 @@ async def refresh_model_maps() -> None:
|
||||
result = await session.exec(query)
|
||||
provider_rows = result.all()
|
||||
|
||||
overrides_by_id: dict[str, tuple[ModelRow, float]] = {}
|
||||
disabled_model_ids: set[str] = set()
|
||||
overrides_by_key: dict[tuple[str, int], tuple[ModelRow, float]] = {}
|
||||
disabled_model_keys: set[tuple[str, int]] = set()
|
||||
|
||||
for provider in provider_rows:
|
||||
if not provider.enabled:
|
||||
continue
|
||||
for model in provider.models:
|
||||
model_key = (model.id.lower(), model.upstream_provider_id)
|
||||
if model.enabled:
|
||||
overrides_by_id[model.id] = (model, provider.provider_fee)
|
||||
overrides_by_key[model_key] = (model, provider.provider_fee)
|
||||
else:
|
||||
disabled_model_ids.add(model.id)
|
||||
disabled_model_keys.add(model_key)
|
||||
|
||||
_model_instances, _provider_map, _unique_models = create_model_mappings(
|
||||
_, _provider_map, _unique_models = create_model_mappings(
|
||||
upstreams=_upstreams,
|
||||
overrides_by_id=overrides_by_id,
|
||||
disabled_model_ids=disabled_model_ids,
|
||||
overrides_by_key=overrides_by_key,
|
||||
disabled_model_keys=disabled_model_keys,
|
||||
)
|
||||
|
||||
|
||||
@@ -166,6 +207,7 @@ _API_PATH_PREFIXES = (
|
||||
"moderations",
|
||||
"providers",
|
||||
"tee/",
|
||||
"attestation",
|
||||
)
|
||||
|
||||
|
||||
@@ -184,20 +226,54 @@ async def proxy(
|
||||
|
||||
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
|
||||
request_body = await request.body()
|
||||
request_body_dict = parse_request_body_json(request_body, path)
|
||||
|
||||
# /tee/* GET requests (e.g. attestation) don't map to models — just
|
||||
# forward to all enabled upstreams without model/cost/auth lookups.
|
||||
if request.method == "GET" and path.startswith("tee/"):
|
||||
all_upstreams = _upstreams
|
||||
# EHBP (Encrypted HTTP Body Protocol) requests carry an Ehbp-Encapsulated-Key
|
||||
# header and a binary HPKE-sealed body. The proxy cannot parse the body to
|
||||
# extract the model id, so the SDK sends it in X-Routstr-Model. Forward the
|
||||
# raw encrypted body to the upstream's /private/ endpoint and stream the
|
||||
# encrypted response back untouched — the SDK's SecureClient decrypts it.
|
||||
is_ehbp = "ehbp-encapsulated-key" in headers
|
||||
if is_ehbp:
|
||||
request_body_dict = {}
|
||||
model_id = headers.get("x-routstr-model", "")
|
||||
if not model_id:
|
||||
return create_error_response(
|
||||
"invalid_request",
|
||||
"EHBP request missing X-Routstr-Model header",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
else:
|
||||
request_body_dict = parse_request_body_json(request_body, path)
|
||||
if is_responses_api:
|
||||
model_id = extract_model_from_responses_request(request_body_dict)
|
||||
else:
|
||||
model_id = request_body_dict.get("model", "unknown")
|
||||
|
||||
# Exact Tinfoil attestation GET routes don't map to models — forward
|
||||
# without model/cost/auth lookups. Do not prefix-match here: paths such as
|
||||
# /attestationjunk must continue through normal authentication.
|
||||
if request.method == "GET" and _is_tinfoil_attestation_path(path):
|
||||
selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams)
|
||||
if not selected_upstreams:
|
||||
return create_error_response(
|
||||
"upstream_error",
|
||||
"No upstream available for unauthenticated GET path",
|
||||
502,
|
||||
request=request,
|
||||
)
|
||||
|
||||
last_error_response = None
|
||||
for i, upstream in enumerate(all_upstreams):
|
||||
for i, upstream in enumerate(selected_upstreams):
|
||||
try:
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
response = await upstream.forward_get_request(request, path, headers)
|
||||
if response.status_code in [502, 429] and i < len(all_upstreams) - 1:
|
||||
if (
|
||||
response.status_code in [502, 429]
|
||||
and i < len(selected_upstreams) - 1
|
||||
):
|
||||
logger.warning(
|
||||
"Upstream %s returned %s for tee GET %s, trying next",
|
||||
"Upstream %s returned %s for unauthenticated GET %s, trying next",
|
||||
upstream.provider_type,
|
||||
response.status_code,
|
||||
path,
|
||||
@@ -206,42 +282,43 @@ async def proxy(
|
||||
return response
|
||||
except UpstreamError as e:
|
||||
logger.warning(
|
||||
"Upstream %s failed for tee GET %s: %s",
|
||||
"Upstream %s failed for unauthenticated GET %s: %s",
|
||||
upstream.provider_type,
|
||||
path,
|
||||
e,
|
||||
)
|
||||
if i == len(all_upstreams) - 1:
|
||||
if i == len(selected_upstreams) - 1:
|
||||
last_error_response = create_upstream_error_response(e, request)
|
||||
continue
|
||||
return last_error_response or create_error_response(
|
||||
"upstream_error", "All upstreams failed", 502, request=request
|
||||
)
|
||||
|
||||
if is_responses_api:
|
||||
model_id = extract_model_from_responses_request(request_body_dict)
|
||||
else:
|
||||
model_id = request_body_dict.get("model", "unknown")
|
||||
candidates = get_candidates(model_id)
|
||||
|
||||
model_obj = get_model_instance(model_id)
|
||||
|
||||
if not model_obj:
|
||||
if not candidates:
|
||||
return create_error_response(
|
||||
"invalid_model", f"Model '{model_id}' not found", 400, request=request
|
||||
)
|
||||
|
||||
upstreams = get_provider_for_model(model_id)
|
||||
if not upstreams:
|
||||
return create_error_response(
|
||||
"invalid_model",
|
||||
f"No provider found for model '{model_id}'",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
if is_ehbp:
|
||||
candidates = [
|
||||
(model, upstream)
|
||||
for model, upstream in candidates
|
||||
if upstream.supports_ehbp
|
||||
]
|
||||
if not candidates:
|
||||
return create_error_response(
|
||||
"unsupported_request",
|
||||
f"No EHBP-capable provider found for model '{model_id}'",
|
||||
400,
|
||||
request=request,
|
||||
)
|
||||
|
||||
# todo figure out cost calculation since fallback provider is usually not the same price
|
||||
# Use first provider for initial checks/cost calculation
|
||||
# primary_upstream = upstreams[0]
|
||||
# Reserve/max-cost checks use the best-ranked candidate; the failover loop
|
||||
# below rebinds (model_obj, upstream) per candidate so forwarding and
|
||||
# settlement always use the model of the provider actually being tried.
|
||||
model_obj = candidates[0][0]
|
||||
|
||||
_max_cost_for_model = await get_max_cost_for_model(
|
||||
model=model_id, session=session, model_obj=model_obj
|
||||
@@ -256,9 +333,25 @@ async def proxy(
|
||||
|
||||
if x_cashu := headers.get("x-cashu", None):
|
||||
last_error = None
|
||||
for i, upstream in enumerate(upstreams):
|
||||
for i, (model_obj, upstream) in enumerate(candidates):
|
||||
try:
|
||||
if is_responses_api:
|
||||
if is_ehbp:
|
||||
if not upstream.supports_ehbp:
|
||||
logger.warning(
|
||||
"Upstream %s does not support EHBP for model=%s",
|
||||
upstream.provider_type,
|
||||
model_id,
|
||||
)
|
||||
continue
|
||||
return await forward_ehbp_x_cashu_request(
|
||||
request=request,
|
||||
x_cashu_token=x_cashu,
|
||||
path=path,
|
||||
max_cost_for_model=max_cost_for_model,
|
||||
model_obj=model_obj,
|
||||
upstream=upstream,
|
||||
)
|
||||
elif is_responses_api:
|
||||
return await upstream.handle_x_cashu_responses(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
)
|
||||
@@ -278,7 +371,7 @@ async def proxy(
|
||||
"status_code": e.status_code,
|
||||
},
|
||||
)
|
||||
if i == len(upstreams) - 1:
|
||||
if i == len(candidates) - 1:
|
||||
last_error = e
|
||||
continue
|
||||
|
||||
@@ -305,12 +398,12 @@ async def proxy(
|
||||
logger.debug("Processing unauthenticated GET request", extra={"path": path})
|
||||
|
||||
last_error_response = None
|
||||
for i, upstream in enumerate(upstreams):
|
||||
for i, (_, upstream) in enumerate(candidates):
|
||||
try:
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
response = await upstream.forward_get_request(request, path, headers)
|
||||
|
||||
if response.status_code in [502, 429] and i < len(upstreams) - 1:
|
||||
if response.status_code in [502, 429] and i < len(candidates) - 1:
|
||||
error_message = ""
|
||||
try:
|
||||
if hasattr(response, "body"):
|
||||
@@ -340,14 +433,14 @@ async def proxy(
|
||||
return response
|
||||
except UpstreamError as e:
|
||||
logger.warning(f"Upstream {upstream.provider_type} failed (GET): {e}")
|
||||
if i == len(upstreams) - 1:
|
||||
if i == len(candidates) - 1:
|
||||
last_error_response = create_upstream_error_response(e, request)
|
||||
continue
|
||||
return last_error_response or create_error_response(
|
||||
"upstream_error", "All upstreams failed", 502, request=request
|
||||
)
|
||||
|
||||
if request_body_dict:
|
||||
if is_ehbp or request_body_dict:
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
|
||||
# Tracks request params already removed in response to upstream rejections,
|
||||
@@ -355,13 +448,59 @@ async def proxy(
|
||||
# the reactive retry can never loop unboundedly.
|
||||
already_stripped: set[str] = set()
|
||||
|
||||
for i, upstream in enumerate(upstreams):
|
||||
for i, (model_obj, upstream) in enumerate(candidates):
|
||||
if i > 0 and request_body_dict:
|
||||
# The reservation was sized to the previous candidate's envelope;
|
||||
# settlement bills the serving candidate, so a pricier fallback
|
||||
# must be re-reserved at its own max cost before it is tried. A
|
||||
# candidate whose envelope the key cannot cover is rejected, just
|
||||
# as it would be had it been ranked first.
|
||||
candidate_max = await get_max_cost_for_model(
|
||||
model=model_id, session=session, model_obj=model_obj
|
||||
)
|
||||
candidate_max = await calculate_discounted_max_cost(
|
||||
candidate_max, request_body_dict, model_obj=model_obj
|
||||
)
|
||||
candidate_max = max(candidate_max, settings.min_request_msat)
|
||||
if candidate_max > max_cost_for_model:
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
try:
|
||||
await pay_for_request(key, candidate_max, session)
|
||||
except HTTPException:
|
||||
if i == len(candidates) - 1:
|
||||
raise
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
continue
|
||||
max_cost_for_model = candidate_max
|
||||
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
if is_responses_api:
|
||||
if is_ehbp:
|
||||
if not upstream.supports_ehbp:
|
||||
logger.warning(
|
||||
"Upstream %s does not support EHBP for model=%s",
|
||||
upstream.provider_type,
|
||||
model_id,
|
||||
)
|
||||
raise UpstreamError(
|
||||
f"Provider {upstream.provider_type} does not support EHBP",
|
||||
status_code=400,
|
||||
)
|
||||
response = await forward_ehbp_request(
|
||||
request=request,
|
||||
path=path,
|
||||
headers=headers,
|
||||
request_body=request_body,
|
||||
upstream=upstream,
|
||||
key=key,
|
||||
max_cost_for_model=max_cost_for_model,
|
||||
session=session,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
elif is_responses_api:
|
||||
response = await upstream.forward_responses_request(
|
||||
request,
|
||||
path,
|
||||
@@ -406,7 +545,7 @@ async def proxy(
|
||||
# When the upstream 400s naming such a param, strip it from the
|
||||
# body and retry the SAME upstream. ``already_stripped`` bounds
|
||||
# this to one retry per distinct param so it always terminates.
|
||||
if response.status_code == 400:
|
||||
if response.status_code == 400 and not is_ehbp:
|
||||
correction = correct_request(
|
||||
request_body,
|
||||
extract_error_message(response),
|
||||
@@ -434,7 +573,7 @@ async def proxy(
|
||||
if response.status_code != 200:
|
||||
# Check if we should retry (502 Upstream Error or 429 Rate Limit)
|
||||
should_retry = response.status_code in [502, 429, 400, 401, 403, 404]
|
||||
if should_retry and i < len(upstreams) - 1:
|
||||
if should_retry and i < len(candidates) - 1:
|
||||
error_message = ""
|
||||
try:
|
||||
if hasattr(response, "body"):
|
||||
@@ -514,12 +653,12 @@ async def proxy(
|
||||
"provider": upstream.provider_type,
|
||||
"model": model_id,
|
||||
"status_code": e.status_code,
|
||||
"retry": i < len(upstreams) - 1,
|
||||
"retry": i < len(candidates) - 1,
|
||||
},
|
||||
)
|
||||
|
||||
# If this was the last provider
|
||||
if i == len(upstreams) - 1:
|
||||
if i == len(candidates) - 1:
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
return create_upstream_error_response(e, request)
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from .openrouter import OpenRouterUpstreamProvider
|
||||
from .perplexity import PerplexityUpstreamProvider
|
||||
from .ppqai import PPQAIUpstreamProvider
|
||||
from .routstr import RoutstrUpstreamProvider
|
||||
from .tinfoil import TinfoilUpstreamProvider
|
||||
from .xai import XAIUpstreamProvider
|
||||
|
||||
upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
|
||||
@@ -26,6 +27,7 @@ upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
|
||||
PerplexityUpstreamProvider,
|
||||
PPQAIUpstreamProvider,
|
||||
RoutstrUpstreamProvider,
|
||||
TinfoilUpstreamProvider,
|
||||
XAIUpstreamProvider,
|
||||
]
|
||||
"""List of all upstream classes"""
|
||||
|
||||
@@ -4,7 +4,14 @@ import json
|
||||
from sqlmodel import select
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.db import UpstreamProviderRow, create_session
|
||||
from ..core.db import (
|
||||
CashuTransaction,
|
||||
UpstreamProviderRow,
|
||||
create_session,
|
||||
)
|
||||
from ..core.db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from ..wallet import send_token
|
||||
from .routstr import RoutstrUpstreamProvider
|
||||
|
||||
@@ -123,7 +130,6 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
||||
},
|
||||
)
|
||||
|
||||
print(amount, mint_url)
|
||||
try:
|
||||
token = await send_token(amount, "sat", mint_url)
|
||||
except Exception as e:
|
||||
@@ -138,6 +144,23 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
await store_cashu_transaction(
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit="sat",
|
||||
mint_url=mint_url,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
)
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Aborting auto top-up because its cashu token could not be persisted",
|
||||
extra={"provider_id": row.id, "mint_url": mint_url},
|
||||
)
|
||||
return
|
||||
|
||||
result = await provider.topup(token)
|
||||
|
||||
if "error" in result:
|
||||
@@ -149,6 +172,26 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
||||
},
|
||||
)
|
||||
else:
|
||||
async with create_session() as session:
|
||||
transaction = (
|
||||
await session.exec(
|
||||
select(CashuTransaction).where(
|
||||
CashuTransaction.token == token,
|
||||
CashuTransaction.type == "out",
|
||||
CashuTransaction.source == "auto_topup",
|
||||
)
|
||||
)
|
||||
).first()
|
||||
if transaction is None:
|
||||
logger.critical(
|
||||
"Completed auto top-up transaction is missing from the database",
|
||||
extra={"provider_id": row.id, "mint_url": mint_url},
|
||||
)
|
||||
else:
|
||||
transaction.collected = True
|
||||
session.add(transaction)
|
||||
await session.commit()
|
||||
|
||||
logger.info(
|
||||
"Auto top-up completed successfully",
|
||||
extra={
|
||||
|
||||
+227
-18
@@ -4,6 +4,7 @@ import asyncio
|
||||
import json
|
||||
import math
|
||||
import traceback
|
||||
import typing
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Iterator
|
||||
from typing import Any, Mapping, Self, cast
|
||||
@@ -20,7 +21,9 @@ from ..core.db import (
|
||||
AsyncSession,
|
||||
UpstreamProviderRow,
|
||||
create_session,
|
||||
store_cashu_transaction,
|
||||
)
|
||||
from ..core.db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from ..core.exceptions import UpstreamError
|
||||
from ..core.redaction import redact_org_ids
|
||||
@@ -55,9 +58,62 @@ from .count_tokens import count_tokens_locally
|
||||
from .litellm_routing import detect_litellm_prefix
|
||||
from .rate_limit import UPSTREAM_RATE_LIMIT, classify_rate_limit
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def _inject_cost_response_headers(
|
||||
headers: dict[str, str], cost_data: CostData | MaxCostData
|
||||
) -> None:
|
||||
"""Inject per-request cost breakdown into response headers.
|
||||
|
||||
The SDK's ``extractUsageFromResponseHeaders`` reads these to populate
|
||||
``inputMsats``, ``outputMsats``, ``totalMsats`` and ``satsCost`` in the
|
||||
usage tracking entry — without them, x-cashu requests show 0.0 for all
|
||||
sat cost fields.
|
||||
"""
|
||||
headers["X-Routstr-Cost-Msats"] = str(cost_data.total_msats)
|
||||
headers["X-Routstr-Input-Cost-Msats"] = str(cost_data.input_msats)
|
||||
headers["X-Routstr-Output-Cost-Msats"] = str(cost_data.output_msats)
|
||||
if cost_data.total_usd:
|
||||
headers["X-Routstr-Cost-Usd"] = str(cost_data.total_usd)
|
||||
|
||||
|
||||
def _inject_cost_into_usage(
|
||||
response_json: dict, cost_data: CostData | MaxCostData
|
||||
) -> None:
|
||||
"""Inject cost breakdown into the response body's ``usage.cost`` object.
|
||||
|
||||
The SDK's ``extractUsageFromResponseBody`` expects ``usage.cost`` to be
|
||||
an object with ``total_msats``/``input_msats``/``output_msats`` (not a
|
||||
plain USD number). When the upstream returns ``cost`` as a number, the
|
||||
SDK cannot extract the msats breakdown from the body alone.
|
||||
"""
|
||||
usage = response_json.get("usage")
|
||||
if not isinstance(usage, dict):
|
||||
return
|
||||
# Direct assignment (not setdefault) so routstr's authoritative cost
|
||||
# data always overwrites any upstream-provided cost values. Using
|
||||
# setdefault would silently keep stale upstream values and drop our
|
||||
# calculated msats breakdown.
|
||||
cost_obj = {
|
||||
"base_msats": cost_data.base_msats,
|
||||
"input_msats": cost_data.input_msats,
|
||||
"output_msats": cost_data.output_msats,
|
||||
"total_msats": cost_data.total_msats,
|
||||
"cache_read_input_tokens": cost_data.cache_read_input_tokens,
|
||||
"cache_creation_input_tokens": cost_data.cache_creation_input_tokens,
|
||||
"cache_read_msats": cost_data.cache_read_msats,
|
||||
"cache_creation_msats": cost_data.cache_creation_msats,
|
||||
}
|
||||
if cost_data.total_usd:
|
||||
cost_obj["total_usd"] = cost_data.total_usd
|
||||
usage["cost"] = cost_obj
|
||||
usage["cost_sats"] = cost_data.total_msats // 1000
|
||||
|
||||
|
||||
def _is_json_content_type(content_type: str | None) -> bool:
|
||||
"""Return True when the upstream response should be parsed as JSON.
|
||||
"""
|
||||
@@ -276,9 +332,27 @@ class BaseUpstreamProvider:
|
||||
|
||||
sats_cost = total_msats // 1000
|
||||
|
||||
# Build the cost object that the SDK's extractUsageFromResponseBody
|
||||
# and extractUsageFromSSEJson expect: an object with total_msats,
|
||||
# input_msats, output_msats, cache_read_msats, cache_creation_msats,
|
||||
# etc. Setting usage.cost to a plain float (total_usd) means the SDK
|
||||
# cannot extract the msats breakdown — cache_read_msats and
|
||||
# cache_creation_msats in particular are lost.
|
||||
cost_obj = {
|
||||
"base_msats": cost_dict.get("base_msats", 0),
|
||||
"input_msats": cost_dict.get("input_msats", 0),
|
||||
"output_msats": cost_dict.get("output_msats", 0),
|
||||
"total_msats": total_msats,
|
||||
"total_usd": total_usd,
|
||||
"cache_read_input_tokens": cost_dict.get("cache_read_input_tokens", 0),
|
||||
"cache_creation_input_tokens": cost_dict.get("cache_creation_input_tokens", 0),
|
||||
"cache_read_msats": cost_dict.get("cache_read_msats", 0),
|
||||
"cache_creation_msats": cost_dict.get("cache_creation_msats", 0),
|
||||
}
|
||||
|
||||
# Inject into top-level usage block (OpenAI/Anthropic style)
|
||||
if "usage" in response_json:
|
||||
response_json["usage"]["cost"] = total_usd
|
||||
response_json["usage"]["cost"] = cost_obj
|
||||
response_json["usage"]["cost_sats"] = sats_cost
|
||||
response_json["usage"]["remaining_balance_msats"] = key.balance
|
||||
self._fold_cache_into_input_tokens(response_json["usage"])
|
||||
@@ -793,6 +867,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
background_tasks: BackgroundTasks,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -835,6 +910,8 @@ class BaseUpstreamProvider:
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
@@ -1007,6 +1084,8 @@ class BaseUpstreamProvider:
|
||||
adjustment_input,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception as e:
|
||||
@@ -1103,6 +1182,7 @@ class BaseUpstreamProvider:
|
||||
session: AsyncSession,
|
||||
deducted_max_cost: int,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> Response:
|
||||
"""Handle non-streaming chat completion responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1149,6 +1229,8 @@ class BaseUpstreamProvider:
|
||||
response_json,
|
||||
session,
|
||||
deducted_max_cost,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
|
||||
await session.refresh(key)
|
||||
@@ -1244,6 +1326,7 @@ class BaseUpstreamProvider:
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming Responses API responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1287,6 +1370,8 @@ class BaseUpstreamProvider:
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
@@ -1416,6 +1501,8 @@ class BaseUpstreamProvider:
|
||||
adjustment_input,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception as e:
|
||||
@@ -1534,6 +1621,7 @@ class BaseUpstreamProvider:
|
||||
session: AsyncSession,
|
||||
deducted_max_cost: int,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> Response:
|
||||
"""Handle non-streaming Responses API responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1583,6 +1671,8 @@ class BaseUpstreamProvider:
|
||||
response_json,
|
||||
session,
|
||||
deducted_max_cost,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
|
||||
await session.refresh(key)
|
||||
@@ -1687,11 +1777,15 @@ class BaseUpstreamProvider:
|
||||
|
||||
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).
|
||||
await adjust_payment_for_tokens(
|
||||
key,
|
||||
{"model": "unknown", "usage": None},
|
||||
session,
|
||||
max_cost,
|
||||
model_obj=None,
|
||||
provider_fee=None,
|
||||
)
|
||||
logger.debug(
|
||||
"Finalized generic streaming payment in background",
|
||||
@@ -1716,6 +1810,7 @@ class BaseUpstreamProvider:
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
async def stream_with_cost(
|
||||
max_cost_for_model: int,
|
||||
@@ -1781,6 +1876,8 @@ class BaseUpstreamProvider:
|
||||
fallback,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
|
||||
@@ -1932,6 +2029,8 @@ class BaseUpstreamProvider:
|
||||
combined_data,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
|
||||
self.inject_cost_metadata(
|
||||
@@ -1979,6 +2078,7 @@ class BaseUpstreamProvider:
|
||||
deducted_max_cost: int,
|
||||
path: str,
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> Response:
|
||||
try:
|
||||
content = await response.aread()
|
||||
@@ -2003,6 +2103,8 @@ class BaseUpstreamProvider:
|
||||
response_json,
|
||||
session,
|
||||
deducted_max_cost,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
|
||||
self.inject_cost_metadata(response_json, cost_data, key)
|
||||
@@ -2026,6 +2128,21 @@ class BaseUpstreamProvider:
|
||||
if k.lower() in allowed_headers
|
||||
}
|
||||
|
||||
# Inject cost breakdown headers so the SDK's
|
||||
# extractUsageFromResponseHeaders can populate
|
||||
# inputMsats/outputMsats/totalMsats for balance-mode requests.
|
||||
if isinstance(cost_data, dict):
|
||||
_cost_data_obj = CostData(
|
||||
base_msats=cost_data.get("base_msats", 0),
|
||||
input_msats=cost_data.get("input_msats", 0),
|
||||
output_msats=cost_data.get("output_msats", 0),
|
||||
total_msats=cost_data.get("total_msats", 0),
|
||||
total_usd=cost_data.get("total_usd", 0.0),
|
||||
)
|
||||
else:
|
||||
_cost_data_obj = cost_data
|
||||
_inject_cost_response_headers(response_headers, _cost_data_obj)
|
||||
|
||||
return Response(
|
||||
content=json.dumps(response_json).encode(),
|
||||
status_code=response.status_code,
|
||||
@@ -2097,6 +2214,7 @@ class BaseUpstreamProvider:
|
||||
key,
|
||||
max_cost_for_model,
|
||||
requested_model,
|
||||
model_obj,
|
||||
)
|
||||
|
||||
response_json = messages_dispatch.coerce_litellm_payload(result)
|
||||
@@ -2108,12 +2226,29 @@ class BaseUpstreamProvider:
|
||||
response_json,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
self.inject_cost_metadata(response_json, cost_data, key)
|
||||
|
||||
# Inject cost breakdown headers for balance-mode requests.
|
||||
if isinstance(cost_data, dict):
|
||||
_cost_data_obj = CostData(
|
||||
base_msats=cost_data.get("base_msats", 0),
|
||||
input_msats=cost_data.get("input_msats", 0),
|
||||
output_msats=cost_data.get("output_msats", 0),
|
||||
total_msats=cost_data.get("total_msats", 0),
|
||||
total_usd=cost_data.get("total_usd", 0.0),
|
||||
)
|
||||
else:
|
||||
_cost_data_obj = cost_data
|
||||
response_headers: dict[str, str] = {}
|
||||
_inject_cost_response_headers(response_headers, _cost_data_obj)
|
||||
|
||||
return Response(
|
||||
content=json.dumps(response_json).encode(),
|
||||
status_code=200,
|
||||
headers=response_headers,
|
||||
media_type="application/json",
|
||||
)
|
||||
|
||||
@@ -2147,6 +2282,7 @@ class BaseUpstreamProvider:
|
||||
requested_model,
|
||||
mint,
|
||||
request_id,
|
||||
model_obj,
|
||||
)
|
||||
|
||||
response_json = messages_dispatch.coerce_litellm_payload(result)
|
||||
@@ -2154,16 +2290,19 @@ class BaseUpstreamProvider:
|
||||
if requested_model and "model" in response_json:
|
||||
response_json["model"] = requested_model
|
||||
|
||||
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
|
||||
cost_data = await self.get_x_cashu_cost(
|
||||
response_json, max_cost_for_model, model_obj
|
||||
)
|
||||
|
||||
if cost_data and "usage" in response_json and isinstance(
|
||||
response_json["usage"], dict
|
||||
):
|
||||
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
|
||||
_inject_cost_into_usage(response_json, cost_data)
|
||||
self._fold_cache_into_input_tokens(response_json["usage"])
|
||||
|
||||
response_headers: dict[str, str] = {}
|
||||
if cost_data:
|
||||
_inject_cost_response_headers(response_headers, cost_data)
|
||||
refund_amount = messages_dispatch.compute_refund(
|
||||
amount, unit, cost_data.total_msats
|
||||
)
|
||||
@@ -2199,6 +2338,7 @@ class BaseUpstreamProvider:
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
requested_model: str | None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Re-emit a litellm Anthropic-event iterator as live SSE bytes
|
||||
with cost reconciliation appended at end of stream."""
|
||||
@@ -2249,6 +2389,8 @@ class BaseUpstreamProvider:
|
||||
fallback,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
usage_finalized = True
|
||||
return (
|
||||
@@ -2322,6 +2464,8 @@ class BaseUpstreamProvider:
|
||||
combined_data,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
)
|
||||
self.inject_cost_metadata(
|
||||
combined_data, cost_data, fresh_key
|
||||
@@ -2359,6 +2503,7 @@ class BaseUpstreamProvider:
|
||||
requested_model: str | None,
|
||||
mint: str | None,
|
||||
request_id: str | None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Buffer a litellm stream end-to-end, compute cost, then replay.
|
||||
|
||||
@@ -2450,7 +2595,7 @@ class BaseUpstreamProvider:
|
||||
}
|
||||
try:
|
||||
cost_data = await self.get_x_cashu_cost(
|
||||
response_data, max_cost_for_model
|
||||
response_data, max_cost_for_model, model_obj
|
||||
)
|
||||
if cost_data:
|
||||
refund_amount = messages_dispatch.compute_refund(
|
||||
@@ -2667,6 +2812,7 @@ class BaseUpstreamProvider:
|
||||
key,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
@@ -2683,6 +2829,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
path,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
finally:
|
||||
await response.aclose()
|
||||
@@ -2698,6 +2845,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
path,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
finally:
|
||||
await response.aclose()
|
||||
@@ -2747,6 +2895,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
background_tasks,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
result.background = background_tasks
|
||||
return result
|
||||
@@ -2760,6 +2909,7 @@ class BaseUpstreamProvider:
|
||||
session,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
finally:
|
||||
await response.aclose()
|
||||
@@ -2845,6 +2995,26 @@ class BaseUpstreamProvider:
|
||||
# Don't revert here — proxy.py owns payment revert to avoid double-revert
|
||||
raise UpstreamError("An unexpected server error occurred", status_code=500)
|
||||
|
||||
supports_ehbp: bool = False
|
||||
|
||||
def get_confidential_inference_profile(self) -> "ConfidentialInferenceProfile | None":
|
||||
"""Return provider policy for encrypted/confidential inference forwarding."""
|
||||
return None
|
||||
|
||||
def get_ehbp_forwarding_target(
|
||||
self, path: str, model_obj: Model
|
||||
) -> "EHBPForwardingTarget":
|
||||
"""Return the EHBP forwarding target for this provider.
|
||||
|
||||
Providers must explicitly opt in by setting ``supports_ehbp = True``
|
||||
and overriding this method. Most upstreams do not accept EHBP-encrypted
|
||||
request bodies, so the base provider intentionally does not provide a
|
||||
default endpoint.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
f"Provider {self.provider_type} does not support EHBP forwarding"
|
||||
)
|
||||
|
||||
async def forward_responses_request(
|
||||
self,
|
||||
request: Request,
|
||||
@@ -2992,6 +3162,7 @@ class BaseUpstreamProvider:
|
||||
key,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
@@ -3007,6 +3178,7 @@ class BaseUpstreamProvider:
|
||||
session,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
finally:
|
||||
await response.aclose()
|
||||
@@ -3183,13 +3355,19 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
|
||||
async def get_x_cashu_cost(
|
||||
self, response_data: dict, max_cost_for_model: int
|
||||
self,
|
||||
response_data: dict,
|
||||
max_cost_for_model: int,
|
||||
model_obj: Model | None,
|
||||
) -> MaxCostData | CostData | None:
|
||||
"""Calculate cost for X-Cashu payment based on response data.
|
||||
|
||||
Args:
|
||||
response_data: Response data containing model and usage information
|
||||
max_cost_for_model: Maximum cost for the model
|
||||
model_obj: The model that actually served the request; billed
|
||||
directly instead of re-deriving pricing from the upstream's
|
||||
echoed model string
|
||||
|
||||
Returns:
|
||||
Cost data object (MaxCostData or CostData) or None if calculation fails
|
||||
@@ -3203,6 +3381,8 @@ class BaseUpstreamProvider:
|
||||
match await calculate_cost(
|
||||
response_data,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
):
|
||||
case MaxCostData() as cost:
|
||||
logger.debug(
|
||||
@@ -3347,6 +3527,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
mint: str | None = None,
|
||||
request_id: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming response for X-Cashu payment, calculating refund if needed.
|
||||
|
||||
@@ -3420,7 +3601,7 @@ class BaseUpstreamProvider:
|
||||
response_data = {"usage": usage_data, "model": model}
|
||||
try:
|
||||
cost_data = await self.get_x_cashu_cost(
|
||||
response_data, max_cost_for_model
|
||||
response_data, max_cost_for_model, model_obj
|
||||
)
|
||||
if cost_data:
|
||||
if unit == "msat":
|
||||
@@ -3469,6 +3650,11 @@ class BaseUpstreamProvider:
|
||||
"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",
|
||||
@@ -3496,9 +3682,7 @@ class BaseUpstreamProvider:
|
||||
and "usage" in data_json
|
||||
and data_json["usage"]
|
||||
):
|
||||
data_json["usage"]["cost_sats"] = (
|
||||
cost_data.total_msats // 1000
|
||||
)
|
||||
_inject_cost_into_usage(data_json, cost_data)
|
||||
changed = True
|
||||
if changed:
|
||||
lines[i] = "data: " + json.dumps(data_json)
|
||||
@@ -3525,6 +3709,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
mint: str | None = None,
|
||||
request_id: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> Response:
|
||||
"""Handle non-streaming response for X-Cashu payment, calculating refund if needed.
|
||||
|
||||
@@ -3546,10 +3731,15 @@ class BaseUpstreamProvider:
|
||||
try:
|
||||
response_json = json.loads(content_str)
|
||||
self._apply_provider_field(response_json)
|
||||
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
|
||||
cost_data = await self.get_x_cashu_cost(
|
||||
response_json, max_cost_for_model, model_obj
|
||||
)
|
||||
|
||||
if cost_data and "usage" in response_json:
|
||||
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
|
||||
# Inject cost breakdown into both the response body (so the
|
||||
# SDK's body extractor picks up the msats breakdown) and the
|
||||
# response headers (so the SDK's header extractor works too).
|
||||
_inject_cost_into_usage(response_json, cost_data)
|
||||
|
||||
if not cost_data:
|
||||
logger.error(
|
||||
@@ -3580,6 +3770,8 @@ class BaseUpstreamProvider:
|
||||
if "content-encoding" in response_headers:
|
||||
del response_headers["content-encoding"]
|
||||
|
||||
_inject_cost_response_headers(response_headers, cost_data)
|
||||
|
||||
if unit == "msat":
|
||||
refund_amount = amount - cost_data.total_msats
|
||||
elif unit == "sat":
|
||||
@@ -3675,6 +3867,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
mint: str | None = None,
|
||||
request_id: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse | Response:
|
||||
"""Handle chat completion response for X-Cashu payment, detecting streaming vs non-streaming.
|
||||
|
||||
@@ -3718,6 +3911,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
mint,
|
||||
request_id=request_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
else:
|
||||
return await self.handle_x_cashu_non_streaming_response(
|
||||
@@ -3728,6 +3922,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
mint,
|
||||
request_id=request_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
@@ -3913,6 +4108,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
mint,
|
||||
request_id=getattr(request.state, "request_id", None),
|
||||
model_obj=model_obj,
|
||||
)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
@@ -4205,6 +4401,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
mint,
|
||||
request_id=getattr(request.state, "request_id", None),
|
||||
model_obj=model_obj,
|
||||
)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
@@ -4255,6 +4452,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
mint: str | None = None,
|
||||
request_id: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse | Response:
|
||||
"""Handle Responses API completion response for X-Cashu payment.
|
||||
|
||||
@@ -4299,6 +4497,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
mint,
|
||||
request_id=request_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
else:
|
||||
return await self.handle_x_cashu_non_streaming_responses_response(
|
||||
@@ -4309,6 +4508,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model,
|
||||
mint,
|
||||
request_id=request_id,
|
||||
model_obj=model_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
@@ -4336,6 +4536,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
mint: str | None = None,
|
||||
request_id: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming Responses API response for X-Cashu payment.
|
||||
|
||||
@@ -4394,7 +4595,7 @@ class BaseUpstreamProvider:
|
||||
response_data = {"usage": usage_data, "model": model}
|
||||
try:
|
||||
cost_data = await self.get_x_cashu_cost(
|
||||
response_data, max_cost_for_model
|
||||
response_data, max_cost_for_model, model_obj
|
||||
)
|
||||
if cost_data:
|
||||
if unit == "msat":
|
||||
@@ -4444,6 +4645,11 @@ class BaseUpstreamProvider:
|
||||
"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",
|
||||
@@ -4471,9 +4677,7 @@ class BaseUpstreamProvider:
|
||||
and "usage" in data_json
|
||||
and data_json["usage"]
|
||||
):
|
||||
data_json["usage"]["cost_sats"] = (
|
||||
cost_data.total_msats // 1000
|
||||
)
|
||||
_inject_cost_into_usage(data_json, cost_data)
|
||||
changed = True
|
||||
if changed:
|
||||
lines[i] = "data: " + json.dumps(data_json)
|
||||
@@ -4500,6 +4704,7 @@ class BaseUpstreamProvider:
|
||||
max_cost_for_model: int,
|
||||
mint: str | None = None,
|
||||
request_id: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
) -> Response:
|
||||
"""Handle non-streaming Responses API response for X-Cashu payment."""
|
||||
logger.debug(
|
||||
@@ -4510,10 +4715,12 @@ class BaseUpstreamProvider:
|
||||
try:
|
||||
response_json = json.loads(content_str)
|
||||
self._apply_provider_field(response_json)
|
||||
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
|
||||
cost_data = await self.get_x_cashu_cost(
|
||||
response_json, max_cost_for_model, model_obj
|
||||
)
|
||||
|
||||
if cost_data and "usage" in response_json:
|
||||
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
|
||||
_inject_cost_into_usage(response_json, cost_data)
|
||||
|
||||
if not cost_data:
|
||||
logger.error(
|
||||
@@ -4544,6 +4751,8 @@ class BaseUpstreamProvider:
|
||||
if "content-encoding" in response_headers:
|
||||
del response_headers["content-encoding"]
|
||||
|
||||
_inject_cost_response_headers(response_headers, cost_data)
|
||||
|
||||
if unit == "msat":
|
||||
refund_amount = amount - cost_data.total_msats
|
||||
elif unit == "sat":
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+20
-10
@@ -94,12 +94,10 @@ async def get_all_models_with_overrides(
|
||||
provider_result = await session.exec(select(UpstreamProviderRow))
|
||||
providers_by_id = {p.id: p for p in provider_result.all()}
|
||||
|
||||
overrides_by_id: dict[str, tuple[ModelRow, float]] = {
|
||||
row.id: (
|
||||
overrides_by_key: dict[tuple[str, int], tuple[ModelRow, float]] = {
|
||||
(row.id.lower(), row.upstream_provider_id): (
|
||||
row,
|
||||
providers_by_id[row.upstream_provider_id].provider_fee
|
||||
if row.upstream_provider_id in providers_by_id
|
||||
else 1.01,
|
||||
providers_by_id[row.upstream_provider_id].provider_fee,
|
||||
)
|
||||
for row in override_rows
|
||||
if row.upstream_provider_id is not None
|
||||
@@ -107,17 +105,28 @@ async def get_all_models_with_overrides(
|
||||
and providers_by_id[row.upstream_provider_id].enabled
|
||||
}
|
||||
|
||||
all_models: dict[str, Model] = {}
|
||||
all_models: dict[tuple[str, str], Model] = {}
|
||||
|
||||
for upstream in upstreams:
|
||||
upstream_db_id = getattr(upstream, "db_id", None)
|
||||
provider_key = (
|
||||
f"db:{upstream_db_id}"
|
||||
if isinstance(upstream_db_id, int)
|
||||
else f"{getattr(upstream, 'provider_type', '')}|{getattr(upstream, 'base_url', '')}"
|
||||
)
|
||||
for model in upstream.get_cached_models():
|
||||
if model.id in overrides_by_id:
|
||||
override_row, provider_fee = overrides_by_id[model.id]
|
||||
all_models[model.id] = _row_to_model(
|
||||
model_key = (
|
||||
(model.id.lower(), upstream_db_id)
|
||||
if isinstance(upstream_db_id, int)
|
||||
else None
|
||||
)
|
||||
if model_key is not None and model_key in overrides_by_key:
|
||||
override_row, provider_fee = overrides_by_key[model_key]
|
||||
all_models[(model.id.lower(), provider_key)] = _row_to_model(
|
||||
override_row, apply_provider_fee=True, provider_fee=provider_fee
|
||||
)
|
||||
elif model.enabled:
|
||||
all_models[model.id] = model
|
||||
all_models[(model.id.lower(), provider_key)] = model
|
||||
|
||||
return list(all_models.values())
|
||||
|
||||
@@ -263,6 +272,7 @@ async def _seed_providers_from_settings(
|
||||
("PERPLEXITY_API_KEY", "perplexity", None, None),
|
||||
("FIREWORKS_API_KEY", "fireworks", None, None),
|
||||
("XAI_API_KEY", "xai", None, None),
|
||||
("TINFOIL_API_KEY", "tinfoil", None, None),
|
||||
]
|
||||
|
||||
for env_key, provider_type, _, _ in env_mappings:
|
||||
|
||||
@@ -8,6 +8,7 @@ from pydantic.v1 import BaseModel, Field
|
||||
from ..core.logging import get_logger
|
||||
from ..payment.models import Architecture, Model, Pricing, async_fetch_openrouter_models
|
||||
from .base import BaseUpstreamProvider, TopupData
|
||||
from .ehbp import EHBPForwardingTarget
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
@@ -39,6 +40,10 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
default_base_url = "https://api.ppq.ai"
|
||||
platform_url = "https://ppq.ai/api-docs"
|
||||
IGNORED_MODEL_IDS: list[str] = ["auto"]
|
||||
# PPQ.AI has a private encrypted endpoint, but this proxy currently has no
|
||||
# provider-attested usage extractor/model binding for it. Keep EHBP disabled
|
||||
# until a ConfidentialInferenceProfile can bill it without max-cost fallback.
|
||||
supports_ehbp = False
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.0):
|
||||
super().__init__(
|
||||
@@ -70,6 +75,20 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
return model_id
|
||||
|
||||
def get_ehbp_forwarding_target(
|
||||
self, path: str, model_obj: Model
|
||||
) -> EHBPForwardingTarget:
|
||||
"""Return the PPQ.AI private enclave target for EHBP requests.
|
||||
|
||||
PPQ.AI exposes EHBP-aware inference under /private/v1/... separate
|
||||
from the public /v1/... endpoint. The encrypted body remains opaque to
|
||||
Routstr, so PPQ.AI also needs X-Private-Model for routing/billing.
|
||||
"""
|
||||
return EHBPForwardingTarget(
|
||||
url=f"{self.base_url.rstrip('/')}/private/{path.lstrip('/')}",
|
||||
headers={"X-Private-Model": model_obj.forwarded_model_id or model_obj.id},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def create_account_static(cls) -> dict[str, object]:
|
||||
"""Create a new PPQ.AI account without requiring an instance.
|
||||
|
||||
@@ -0,0 +1,236 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
from pydantic.v1 import BaseModel
|
||||
|
||||
from ..core.exceptions import UpstreamError
|
||||
from ..core.logging import get_logger
|
||||
from ..payment.models import Architecture, Model, Pricing
|
||||
from .base import BaseUpstreamProvider
|
||||
from .ehbp import (
|
||||
_ENCLAVE_URL_HEADER,
|
||||
_PROXY_ONLY_HEADERS,
|
||||
_RESPONSE_USAGE_HEADER,
|
||||
ConfidentialInferenceProfile,
|
||||
EHBPForwardingTarget,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class TinfoilModelPricing(BaseModel):
|
||||
inputTokenPricePer1M: float = 0.0
|
||||
outputTokenPricePer1M: float = 0.0
|
||||
requestPrice: float = 0.0
|
||||
|
||||
|
||||
class TinfoilModel(BaseModel):
|
||||
id: str
|
||||
context_window: int = 0
|
||||
created: int = 0
|
||||
multimodal: bool = False
|
||||
reasoning: bool = False
|
||||
tool_calling: bool = False
|
||||
type: str = "chat"
|
||||
pricing: TinfoilModelPricing = TinfoilModelPricing()
|
||||
endpoints: list[str] = []
|
||||
|
||||
|
||||
class TinfoilUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Direct upstream provider for the Tinfoil inference API.
|
||||
|
||||
Tinfoil hosts open-source models inside attested secure enclaves and exposes
|
||||
an OpenAI-compatible API at ``https://inference.tinfoil.sh``. Request and
|
||||
response bodies are encrypted end-to-end with EHBP (HPKE), so Routstr acts
|
||||
as a blind relay: it forwards the opaque encrypted body, never sees
|
||||
plaintext, and bills from the ``X-Tinfoil-Usage-Metrics`` header that
|
||||
Tinfoil returns outside the encrypted body when
|
||||
``X-Tinfoil-Request-Usage-Metrics: true`` is set.
|
||||
"""
|
||||
|
||||
provider_type = "tinfoil"
|
||||
default_base_url = "https://inference.tinfoil.sh"
|
||||
platform_url = "https://docs.tinfoil.sh"
|
||||
supports_ehbp = True
|
||||
confidential_inference_profile = ConfidentialInferenceProfile(
|
||||
usage_response_header=_RESPONSE_USAGE_HEADER,
|
||||
client_target_url_header=_ENCLAVE_URL_HEADER,
|
||||
allow_client_target_override=True,
|
||||
proxy_only_headers=_PROXY_ONLY_HEADERS,
|
||||
)
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.0):
|
||||
super().__init__(
|
||||
base_url=self.default_base_url,
|
||||
api_key=api_key,
|
||||
provider_fee=provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _build_from_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "TinfoilUpstreamProvider":
|
||||
return cls(
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Tinfoil",
|
||||
"default_base_url": cls.default_base_url,
|
||||
"fixed_base_url": True,
|
||||
"platform_url": cls.platform_url,
|
||||
"can_create_account": False,
|
||||
"can_topup": False,
|
||||
"can_show_balance": False,
|
||||
}
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
return model_id.removeprefix("tinfoil/")
|
||||
|
||||
def get_confidential_inference_profile(self) -> ConfidentialInferenceProfile:
|
||||
return self.confidential_inference_profile
|
||||
|
||||
async def forward_get_request(
|
||||
self,
|
||||
request: Request,
|
||||
path: str,
|
||||
headers: dict,
|
||||
) -> Response | StreamingResponse:
|
||||
"""Handle Tinfoil-specific GET endpoints.
|
||||
|
||||
* ``/attestation`` (or ``/tee/attestation``): proxy to the Tinfoil ATC
|
||||
(attestation bundle proxy) at ``https://atc.tinfoil.sh/attestation``.
|
||||
* Other GETs: forward to the provider base URL
|
||||
(``https://inference.tinfoil.sh``). ``X-Tinfoil-Enclave-Url`` is an
|
||||
EHBP-only header used for encrypted POST requests and is not honored
|
||||
for unencrypted GET requests.
|
||||
"""
|
||||
clean_path = path.removeprefix("tee/").rstrip("/")
|
||||
if clean_path == "attestation":
|
||||
return await self._proxy_attestation(headers)
|
||||
return await super().forward_get_request(request, path, headers)
|
||||
|
||||
async def _proxy_attestation(self, headers: dict) -> Response:
|
||||
url = "https://atc.tinfoil.sh/attestation"
|
||||
async with httpx.AsyncClient(
|
||||
transport=httpx.AsyncHTTPTransport(retries=1),
|
||||
timeout=30.0,
|
||||
) as client:
|
||||
try:
|
||||
resp = await client.get(
|
||||
url,
|
||||
headers={
|
||||
"Accept": headers.get("accept", "application/json"),
|
||||
},
|
||||
)
|
||||
response_headers = dict(resp.headers)
|
||||
response_headers.pop("content-encoding", None)
|
||||
response_headers.pop("content-length", None)
|
||||
return Response(
|
||||
content=resp.content,
|
||||
status_code=resp.status_code,
|
||||
headers=response_headers,
|
||||
)
|
||||
except Exception as exc:
|
||||
raise UpstreamError(
|
||||
f"Error fetching Tinfoil attestation: {type(exc).__name__}",
|
||||
status_code=502,
|
||||
) from exc
|
||||
|
||||
def get_ehbp_forwarding_target(
|
||||
self, path: str, model_obj: Model
|
||||
) -> EHBPForwardingTarget:
|
||||
"""Return the Tinfoil enclave target for EHBP requests.
|
||||
|
||||
Requests usage metrics from the enclave so Routstr can bill exactly
|
||||
without decrypting the response body. The actual forwarding URL is
|
||||
overridden at dispatch time by ``X-Tinfoil-Enclave-Url`` when the SDK
|
||||
sends it (see ``routstr/upstream/ehbp.py``).
|
||||
"""
|
||||
return EHBPForwardingTarget(
|
||||
url=f"{self.base_url.rstrip('/')}/{path.lstrip('/')}",
|
||||
headers={"X-Tinfoil-Request-Usage-Metrics": "true"},
|
||||
profile=self.confidential_inference_profile,
|
||||
)
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch models from the public Tinfoil models endpoint.
|
||||
|
||||
``GET /v1/models`` is unauthenticated and returns all available models
|
||||
with their pricing in USD per 1M tokens.
|
||||
"""
|
||||
url = f"{self.base_url}/v1/models"
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.get(url)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
models_data = data.get("data", [])
|
||||
|
||||
models: list[Model] = []
|
||||
for model_data in models_data:
|
||||
try:
|
||||
tf = TinfoilModel.parse_obj(model_data)
|
||||
input_price = tf.pricing.inputTokenPricePer1M
|
||||
output_price = tf.pricing.outputTokenPricePer1M
|
||||
request_price = tf.pricing.requestPrice
|
||||
|
||||
modality = "text->text"
|
||||
input_modalities = ["text"]
|
||||
output_modalities = ["text"]
|
||||
if tf.multimodal:
|
||||
modality = "text->text+image"
|
||||
input_modalities = ["text", "image"]
|
||||
|
||||
models.append(
|
||||
Model(
|
||||
id=tf.id,
|
||||
name=tf.id,
|
||||
created=tf.created,
|
||||
description=f"Tinfoil {tf.type} model",
|
||||
context_length=tf.context_window,
|
||||
architecture=Architecture(
|
||||
modality=modality,
|
||||
input_modalities=input_modalities,
|
||||
output_modalities=output_modalities,
|
||||
tokenizer="Unknown",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(
|
||||
prompt=input_price / 1_000_000,
|
||||
completion=output_price / 1_000_000,
|
||||
request=request_price,
|
||||
image=0.0,
|
||||
web_search=0.0,
|
||||
internal_reasoning=0.0,
|
||||
),
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to parse Tinfoil model",
|
||||
extra={
|
||||
"model_id": model_data.get("id", "unknown"),
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
|
||||
return models
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error fetching models from Tinfoil",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
return []
|
||||
@@ -0,0 +1,189 @@
|
||||
"""h11-based HTTP client for EHBP requests that captures HTTP trailers.
|
||||
|
||||
httpx/httpcore silently discard HTTP trailers during chunked transfer
|
||||
decoding. Tinfoil returns ``X-Tinfoil-Usage-Metrics`` as a trailer on
|
||||
streaming responses, so we need a lower-level HTTP client that preserves
|
||||
trailers from the h11 ``EndOfMessage`` event.
|
||||
|
||||
Because EHBP response bodies are opaque encrypted blobs, buffering the full
|
||||
response is acceptable — the client decrypts the complete body regardless of
|
||||
whether it arrived streamed or buffered.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ssl
|
||||
from dataclasses import dataclass, field
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import h11
|
||||
|
||||
from ..core import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
_READ_BUFSIZE = 65536
|
||||
_DEFAULT_TIMEOUT_SECONDS = 30.0
|
||||
_DEFAULT_CLOSE_TIMEOUT_SECONDS = 1.0
|
||||
_DEFAULT_MAX_RESPONSE_BYTES = 25 * 1024 * 1024
|
||||
_HOP_BY_HOP_HEADERS = {
|
||||
"connection",
|
||||
"keep-alive",
|
||||
"proxy-authenticate",
|
||||
"proxy-authorization",
|
||||
"te",
|
||||
"trailer",
|
||||
"transfer-encoding",
|
||||
"upgrade",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrailerResponse:
|
||||
"""Buffered HTTP response with optional trailer headers."""
|
||||
|
||||
status_code: int
|
||||
headers: list[tuple[str, str]]
|
||||
body: bytes
|
||||
trailers: list[tuple[str, str]] = field(default_factory=list)
|
||||
|
||||
|
||||
def _get_header(headers: list[tuple[str, str]], name: str) -> str | None:
|
||||
name_lower = name.lower()
|
||||
for k, v in headers:
|
||||
if k.lower() == name_lower:
|
||||
return v
|
||||
return None
|
||||
|
||||
|
||||
def _strip_hop_by_hop_headers(headers: dict[str, str]) -> dict[str, str]:
|
||||
"""Remove connection-specific headers before serializing a new request."""
|
||||
connection_tokens: set[str] = set()
|
||||
for key, value in headers.items():
|
||||
if key.lower() == "connection":
|
||||
connection_tokens.update(
|
||||
token.strip().lower() for token in value.split(",") if token.strip()
|
||||
)
|
||||
|
||||
excluded = _HOP_BY_HOP_HEADERS | connection_tokens
|
||||
return {key: value for key, value in headers.items() if key.lower() not in excluded}
|
||||
|
||||
|
||||
async def forward_with_trailer(
|
||||
*,
|
||||
method: str,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
body: bytes,
|
||||
timeout_seconds: float = _DEFAULT_TIMEOUT_SECONDS,
|
||||
max_response_bytes: int = _DEFAULT_MAX_RESPONSE_BYTES,
|
||||
close_timeout_seconds: float = _DEFAULT_CLOSE_TIMEOUT_SECONDS,
|
||||
) -> TrailerResponse:
|
||||
"""Send an HTTP/1.1 request via h11 and capture HTTP trailers.
|
||||
|
||||
Returns a :class:`TrailerResponse` with the full buffered body and any
|
||||
trailer headers from the ``EndOfMessage`` event.
|
||||
"""
|
||||
parsed = urlsplit(url)
|
||||
host = parsed.hostname
|
||||
if not host:
|
||||
raise ValueError(f"Invalid URL (no hostname): {url}")
|
||||
port = parsed.port or 443
|
||||
path = parsed.path or "/"
|
||||
if parsed.query:
|
||||
path = f"{path}?{parsed.query}"
|
||||
|
||||
# FastAPI has already decoded the incoming request body. Do not carry the
|
||||
# original connection's framing or other hop-by-hop metadata into the new
|
||||
# upstream connection.
|
||||
headers = _strip_hop_by_hop_headers(headers)
|
||||
|
||||
ssl_ctx = ssl.create_default_context()
|
||||
reader, writer = await asyncio.wait_for(
|
||||
asyncio.open_connection(host, port, ssl=ssl_ctx),
|
||||
timeout=timeout_seconds,
|
||||
)
|
||||
|
||||
try:
|
||||
# Build HTTP/1.1 request
|
||||
header_lines = [f"{method} {path} HTTP/1.1"]
|
||||
has_host = any(k.lower() == "host" for k in headers)
|
||||
if not has_host:
|
||||
header_lines.append(f"Host: {host}")
|
||||
header_lines.append("Connection: close")
|
||||
|
||||
for key, value in headers.items():
|
||||
if key.lower() == "host":
|
||||
continue
|
||||
header_lines.append(f"{key}: {value}")
|
||||
|
||||
if body and not any(k.lower() == "content-length" for k in headers):
|
||||
header_lines.append(f"Content-Length: {len(body)}")
|
||||
|
||||
request_data = "\r\n".join(header_lines).encode() + b"\r\n\r\n"
|
||||
if body:
|
||||
request_data += body
|
||||
|
||||
writer.write(request_data)
|
||||
await asyncio.wait_for(writer.drain(), timeout=timeout_seconds)
|
||||
|
||||
# Parse response with h11
|
||||
conn = h11.Connection(h11.CLIENT)
|
||||
status_code = 0
|
||||
resp_headers: list[tuple[str, str]] = []
|
||||
body_chunks: list[bytes] = []
|
||||
body_size = 0
|
||||
trailers: list[tuple[str, str]] = []
|
||||
|
||||
while True:
|
||||
event = conn.next_event()
|
||||
|
||||
if event is h11.NEED_DATA:
|
||||
data = await asyncio.wait_for(
|
||||
reader.read(_READ_BUFSIZE),
|
||||
timeout=timeout_seconds,
|
||||
)
|
||||
conn.receive_data(data if data else b"")
|
||||
continue
|
||||
|
||||
if isinstance(event, h11.Response):
|
||||
status_code = event.status_code
|
||||
resp_headers = [(k.decode(), v.decode()) for k, v in event.headers]
|
||||
|
||||
elif isinstance(event, h11.Data):
|
||||
body_size += len(event.data)
|
||||
if body_size > max_response_bytes:
|
||||
raise ValueError(
|
||||
f"EHBP response exceeded {max_response_bytes} bytes"
|
||||
)
|
||||
body_chunks.append(event.data)
|
||||
|
||||
elif isinstance(event, h11.EndOfMessage):
|
||||
for k, v in event.headers:
|
||||
trailers.append((k.decode(), v.decode()))
|
||||
break
|
||||
|
||||
elif isinstance(event, h11.PAUSED):
|
||||
# Shouldn't happen for simple request/response, but break safely
|
||||
logger.warning("h11 PAUSED event during EHBP response parsing")
|
||||
break
|
||||
|
||||
elif isinstance(event, h11.ConnectionClosed):
|
||||
break
|
||||
|
||||
return TrailerResponse(
|
||||
status_code=status_code,
|
||||
headers=resp_headers,
|
||||
body=b"".join(body_chunks),
|
||||
trailers=trailers,
|
||||
)
|
||||
finally:
|
||||
writer.close()
|
||||
if close_timeout_seconds > 0:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
writer.wait_closed(), timeout=close_timeout_seconds
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
+168
-96
@@ -14,7 +14,7 @@ from pydantic_core import PydanticUndefined
|
||||
from sqlmodel import col, select, update
|
||||
|
||||
from .core import db, get_logger
|
||||
from .core.db import store_cashu_transaction
|
||||
from .core.db import store_cashu_transaction_with_retry as store_cashu_transaction
|
||||
from .core.settings import settings
|
||||
from .payment.lnurl import raw_send_to_lnurl
|
||||
|
||||
@@ -219,8 +219,12 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
|
||||
"""Internal send function - returns amount and serialized token"""
|
||||
effective_mint_url = mint_url or settings.primary_mint
|
||||
wallet: Wallet = await get_wallet(effective_mint_url, unit)
|
||||
proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
|
||||
all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
|
||||
proofs = [proof for proof in all_proofs if not proof.reserved]
|
||||
# Fallback must compare the requested amount with liquid proofs only. Counting
|
||||
# reserved proofs here can suppress fallback even though they cannot be sent.
|
||||
proofs_for_mint = sum(p.amount for p in proofs)
|
||||
reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved)
|
||||
|
||||
# Fallback: proofs from untrusted source mints are swapped to primary_mint
|
||||
# during receive, so the user's preferred refund_mint_url may have no proofs
|
||||
@@ -232,8 +236,10 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
|
||||
)
|
||||
effective_mint_url = settings.primary_mint
|
||||
wallet = await get_wallet(effective_mint_url, unit)
|
||||
proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
|
||||
all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
|
||||
proofs = [proof for proof in all_proofs if not proof.reserved]
|
||||
proofs_for_mint = sum(p.amount for p in proofs)
|
||||
reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved)
|
||||
|
||||
all_mint_urls = list({k.mint_url for k in wallet.keysets.values()})
|
||||
proof_summary = {
|
||||
@@ -247,7 +253,8 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
|
||||
raw_proofs_by_keyset[p.id] = raw_proofs_by_keyset.get(p.id, 0) + p.amount
|
||||
logger.info(
|
||||
f"send: proof inventory | mint={effective_mint_url} unit={unit} amount={amount} "
|
||||
f"primary_mint={settings.primary_mint} proofs_for_mint={proofs_for_mint} "
|
||||
f"primary_mint={settings.primary_mint} liquid_proofs_for_mint={proofs_for_mint} "
|
||||
f"reserved_proofs_for_mint={reserved_for_mint} "
|
||||
f"all_mints={all_mint_urls} by_keyset={proof_summary} "
|
||||
f"raw_proofs_by_keyset_id={raw_proofs_by_keyset} "
|
||||
f"total_wallet_proofs={sum(p.amount for p in wallet.proofs)}"
|
||||
@@ -745,11 +752,11 @@ async def credit_balance(
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logger.debug(
|
||||
"Cashu token successfully redeemed and stored",
|
||||
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"Cashu token successfully redeemed and stored",
|
||||
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
|
||||
)
|
||||
return amount
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
@@ -848,7 +855,7 @@ async def fetch_all_balances(
|
||||
"unit": unit,
|
||||
"wallet_balance": proofs_balance,
|
||||
"user_balance": user_balance,
|
||||
"owner_balance": proofs_balance - user_balance if proofs_balance != 0 else 0,
|
||||
"owner_balance": proofs_balance - user_balance,
|
||||
}
|
||||
return result
|
||||
except Exception as e:
|
||||
@@ -904,9 +911,7 @@ async def fetch_all_balances(
|
||||
total_wallet_balance_sats += proofs_balance_sats
|
||||
total_user_balance_sats += user_balance_sats
|
||||
|
||||
owner_balance = 0
|
||||
if total_wallet_balance_sats != 0:
|
||||
owner_balance = total_wallet_balance_sats - total_user_balance_sats
|
||||
owner_balance = total_wallet_balance_sats - total_user_balance_sats
|
||||
|
||||
return (
|
||||
balance_details,
|
||||
@@ -919,104 +924,132 @@ async def fetch_all_balances(
|
||||
async def periodic_payout() -> None:
|
||||
while True:
|
||||
await asyncio.sleep(settings.payout_interval_seconds)
|
||||
print(settings.payout_interval_seconds)
|
||||
if not settings.receive_ln_address:
|
||||
continue
|
||||
try:
|
||||
if not settings.receive_ln_address:
|
||||
continue
|
||||
|
||||
# Include the primary mint even if it is not listed in cashu_mints,
|
||||
# matching fetch_all_balances(); otherwise primary-mint funds never
|
||||
# auto-payout.
|
||||
mint_urls: list[str] = list(settings.cashu_mints)
|
||||
if settings.primary_mint and settings.primary_mint not in mint_urls:
|
||||
mint_urls.append(settings.primary_mint)
|
||||
|
||||
async with db.create_session() as session:
|
||||
for mint_url in settings.cashu_mints:
|
||||
for mint_url in mint_urls:
|
||||
for unit in ["sat", "msat"]:
|
||||
wallet = await get_wallet(mint_url, unit)
|
||||
proofs = get_proofs_per_mint_and_unit(
|
||||
wallet, mint_url, unit, not_reserved=True
|
||||
)
|
||||
proofs = await slow_filter_spend_proofs(proofs, wallet)
|
||||
await asyncio.sleep(5)
|
||||
user_balance = await db.balances_for_mint_and_unit(
|
||||
session, mint_url, unit
|
||||
)
|
||||
if unit == "sat":
|
||||
user_balance = user_balance // 1000
|
||||
proofs_balance = sum(proof.amount for proof in proofs)
|
||||
available_balance = proofs_balance - user_balance
|
||||
# Threshold is configured in sats; convert for msat wallets.
|
||||
min_amount = (
|
||||
settings.min_payout_sat
|
||||
if unit == "sat"
|
||||
else settings.min_payout_sat * 1000
|
||||
)
|
||||
if available_balance > min_amount:
|
||||
amount_received = await raw_send_to_lnurl(
|
||||
wallet,
|
||||
proofs,
|
||||
settings.receive_ln_address,
|
||||
unit,
|
||||
amount=available_balance,
|
||||
# Isolate failures per mint/unit so one slow or failing
|
||||
# mint does not abort payout for every other mint/unit.
|
||||
try:
|
||||
wallet = await get_wallet(mint_url, unit)
|
||||
proofs = get_proofs_per_mint_and_unit(
|
||||
wallet, mint_url, unit, not_reserved=True
|
||||
)
|
||||
logger.info(
|
||||
"Payout sent successfully",
|
||||
proofs = await slow_filter_spend_proofs(proofs, wallet)
|
||||
await asyncio.sleep(5)
|
||||
user_balance = await db.balances_for_mint_and_unit(
|
||||
session, mint_url, unit
|
||||
)
|
||||
if unit == "sat":
|
||||
user_balance = user_balance // 1000
|
||||
proofs_balance = sum(proof.amount for proof in proofs)
|
||||
available_balance = proofs_balance - user_balance
|
||||
# Threshold is configured in sats; convert for msat wallets.
|
||||
min_amount = (
|
||||
settings.min_payout_sat
|
||||
if unit == "sat"
|
||||
else settings.min_payout_sat * 1000
|
||||
)
|
||||
if available_balance > min_amount:
|
||||
amount_received = await raw_send_to_lnurl(
|
||||
wallet,
|
||||
proofs,
|
||||
settings.receive_ln_address,
|
||||
unit,
|
||||
amount=available_balance,
|
||||
)
|
||||
logger.info(
|
||||
"Payout sent successfully",
|
||||
extra={
|
||||
"mint_url": mint_url,
|
||||
"unit": unit,
|
||||
"balance": available_balance,
|
||||
"amount_received": amount_received,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error sending payout: {type(e).__name__}",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"mint_url": mint_url,
|
||||
"unit": unit,
|
||||
"balance": available_balance,
|
||||
"amount_received": amount_received,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error sending payout: {type(e).__name__}",
|
||||
f"Error in periodic payout cycle: {type(e).__name__}",
|
||||
extra={"error": str(e)},
|
||||
)
|
||||
|
||||
|
||||
async def _refund_sweep_once(cutoff: int) -> None:
|
||||
async with db.create_session() as session:
|
||||
stmt = select(db.CashuTransaction).where(
|
||||
db.CashuTransaction.type == "out",
|
||||
db.CashuTransaction.collected == False, # noqa: E712
|
||||
db.CashuTransaction.swept == False, # noqa: E712
|
||||
db.CashuTransaction.created_at < cutoff,
|
||||
)
|
||||
results = await session.exec(stmt)
|
||||
refunds = results.all()
|
||||
|
||||
for refund in refunds:
|
||||
try:
|
||||
await recieve_token(refund.token)
|
||||
refund.swept = True
|
||||
session.add(refund)
|
||||
logger.info(
|
||||
"Swept uncollected refund",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
"amount": refund.amount,
|
||||
"unit": refund.unit,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
error_msg = str(e).lower()
|
||||
if "already spent" in error_msg:
|
||||
refund.collected = True
|
||||
session.add(refund)
|
||||
logger.info(
|
||||
"Refund already spent (client collected), marking swept",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Failed to sweep refund",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
"error": str(e),
|
||||
},
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def refund_sweep_once() -> None:
|
||||
"""Sweep eligible uncollected refund tokens once."""
|
||||
cutoff = int(time.time()) - settings.refund_sweep_ttl_seconds
|
||||
await _refund_sweep_once(cutoff)
|
||||
|
||||
|
||||
async def periodic_refund_sweep() -> None:
|
||||
while True:
|
||||
await asyncio.sleep(60 * 60) # every hour
|
||||
try:
|
||||
cutoff = int(time.time()) - settings.refund_sweep_ttl_seconds
|
||||
async with db.create_session() as session:
|
||||
stmt = select(db.CashuTransaction).where(
|
||||
db.CashuTransaction.type == "out",
|
||||
db.CashuTransaction.collected == False, # noqa: E712
|
||||
db.CashuTransaction.swept == False, # noqa: E712
|
||||
db.CashuTransaction.created_at < cutoff,
|
||||
)
|
||||
results = await session.exec(stmt)
|
||||
refunds = results.all()
|
||||
|
||||
for refund in refunds:
|
||||
try:
|
||||
await recieve_token(refund.token)
|
||||
refund.swept = True
|
||||
session.add(refund)
|
||||
logger.info(
|
||||
"Swept uncollected refund",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
"amount": refund.amount,
|
||||
"unit": refund.unit,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
error_msg = str(e).lower()
|
||||
if "already spent" in error_msg:
|
||||
refund.collected = True
|
||||
session.add(refund)
|
||||
logger.info(
|
||||
"Refund already spent (client collected), marking swept",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Failed to sweep refund",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
"error": str(e),
|
||||
},
|
||||
)
|
||||
await session.commit()
|
||||
await refund_sweep_once()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error in periodic refund sweep",
|
||||
@@ -1039,17 +1072,56 @@ async def periodic_routstr_fee_payout() -> None:
|
||||
try:
|
||||
async with db.create_session() as session:
|
||||
fee = await db.get_routstr_fee(session)
|
||||
if fee.payout_in_progress_msats:
|
||||
logger.critical(
|
||||
"Routstr fee payout requires manual reconciliation",
|
||||
extra={
|
||||
"payout_in_progress_msats": fee.payout_in_progress_msats,
|
||||
"payout_started_at": fee.payout_started_at,
|
||||
},
|
||||
)
|
||||
continue
|
||||
|
||||
accumulated_sats = fee.accumulated_msats // 1000
|
||||
if accumulated_sats >= ROUTSTR_FEE_DEFAULT_PAYOUT:
|
||||
wallet = await get_wallet(settings.primary_mint, "sat")
|
||||
proofs = get_proofs_per_mint_and_unit(
|
||||
wallet, settings.primary_mint, "sat", not_reserved=True
|
||||
)
|
||||
amount_received = await raw_send_to_lnurl(
|
||||
wallet, proofs, ROUTSTR_LN_ADDRESS, "sat", amount=accumulated_sats
|
||||
)
|
||||
paid_msats = accumulated_sats * 1000
|
||||
await db.reset_routstr_fee(session, paid_msats)
|
||||
payout_checkpointed = await db.reset_routstr_fee(
|
||||
session, paid_msats
|
||||
)
|
||||
if not payout_checkpointed:
|
||||
logger.warning("Routstr fee payout was already claimed")
|
||||
continue
|
||||
|
||||
try:
|
||||
amount_received = await raw_send_to_lnurl(
|
||||
wallet,
|
||||
proofs,
|
||||
ROUTSTR_LN_ADDRESS,
|
||||
"sat",
|
||||
amount=accumulated_sats,
|
||||
)
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Routstr fee payout outcome is unknown; manual reconciliation required",
|
||||
extra={"payout_in_progress_msats": paid_msats},
|
||||
exc_info=True,
|
||||
)
|
||||
continue
|
||||
|
||||
payout_completed = await db.complete_routstr_fee_payout(
|
||||
session, paid_msats
|
||||
)
|
||||
if not payout_completed:
|
||||
logger.critical(
|
||||
"Routstr fee payout sent but checkpoint was not completed",
|
||||
extra={"payout_in_progress_msats": paid_msats},
|
||||
)
|
||||
continue
|
||||
|
||||
logger.info(
|
||||
"Routstr fee payout sent",
|
||||
extra={
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Operational and recovery scripts for routstr-core (not a shipped package)."""
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Recover admin access by resetting the stored admin password (issue #553).
|
||||
|
||||
The lockout escape hatch for an operator who has lost the admin password. It
|
||||
talks to the ``secrets`` table directly and deliberately does *not* require
|
||||
``ROUTSTR_SECRET_KEY``: the admin password is scrypt-hashed (key-independent),
|
||||
so recovery works even when the encryption key is missing or has changed.
|
||||
|
||||
Two explicit, mutually exclusive actions — running with no arguments only prints
|
||||
help, so the password can't be reset by accident:
|
||||
|
||||
python scripts/reset_admin_password.py --password <new-password>
|
||||
Hash and store <new-password> now.
|
||||
|
||||
python scripts/reset_admin_password.py --regenerate
|
||||
Clear the stored hash; the next node startup generates a fresh random
|
||||
password and logs it once (with the /admin URL).
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import sys
|
||||
import time
|
||||
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import create_session, get_secret, set_admin_password
|
||||
from routstr.core.vault import MIN_PASSWORD_LENGTH
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="reset_admin_password",
|
||||
description="Reset the node's admin password (recovery from lockout).",
|
||||
)
|
||||
action = parser.add_mutually_exclusive_group()
|
||||
action.add_argument(
|
||||
"--password",
|
||||
metavar="NEW_PASSWORD",
|
||||
help=f"set this as the new admin password (min {MIN_PASSWORD_LENGTH} chars)",
|
||||
)
|
||||
action.add_argument(
|
||||
"--regenerate",
|
||||
action="store_true",
|
||||
help="clear the password so the next startup generates and logs a new one",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
async def apply_reset(
|
||||
session: AsyncSession,
|
||||
*,
|
||||
password: str | None = None,
|
||||
regenerate: bool = False,
|
||||
) -> str:
|
||||
"""Perform the requested reset against ``session``; return a status message."""
|
||||
if password is not None:
|
||||
if len(password) < MIN_PASSWORD_LENGTH:
|
||||
raise ValueError(
|
||||
f"New password must be at least {MIN_PASSWORD_LENGTH} characters"
|
||||
)
|
||||
await set_admin_password(session, password)
|
||||
return "Admin password updated."
|
||||
|
||||
if regenerate:
|
||||
secret = await get_secret(session)
|
||||
secret.admin_password_hash = None
|
||||
secret.updated_at = int(time.time())
|
||||
session.add(secret)
|
||||
await session.commit()
|
||||
return (
|
||||
"Admin password cleared. The next node startup will generate a new "
|
||||
"one and log it once with the /admin URL."
|
||||
)
|
||||
|
||||
return ""
|
||||
|
||||
|
||||
async def _run(password: str | None, regenerate: bool) -> str:
|
||||
async with create_session() as session:
|
||||
return await apply_reset(
|
||||
session, password=password, regenerate=regenerate
|
||||
)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
parser = build_parser()
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
if args.password is None and not args.regenerate:
|
||||
parser.print_help()
|
||||
return 0
|
||||
|
||||
try:
|
||||
message = asyncio.run(_run(args.password, args.regenerate))
|
||||
except ValueError as exc:
|
||||
print(f"error: {exc}", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
print(message)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,15 @@
|
||||
"""Shared pytest configuration for the whole suite.
|
||||
|
||||
A fixed, valid ``ROUTSTR_SECRET_KEY`` is set before any app import so that
|
||||
secret encryption is deterministic across the suite and the mandatory-key
|
||||
fail-fast does not break app-boot tests. Tests that need a different key (or an
|
||||
absent one) override this per-test via ``monkeypatch``.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
# Valid Fernet keys; KEY_A is the suite default, KEY_B is for wrong-key tests.
|
||||
TEST_SECRET_KEY = "l_Tkp-7xmjcQ-IFhr6qhILrU8HPRbEmYMrfSbo_5srU="
|
||||
TEST_SECRET_KEY_ALT = "_Teyrky_iToeDK51Tj1FsI9MJ340_cqKGmeher-a7MQ="
|
||||
|
||||
os.environ.setdefault("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY)
|
||||
@@ -0,0 +1,121 @@
|
||||
"""Tests for admin password auth backed by the hashed Secret store (issue #553).
|
||||
|
||||
Login and password change verify against the one-way ``Secret.admin_password_hash``
|
||||
(scrypt) instead of a plaintext settings field, which also closes the old ``!=``
|
||||
timing-attack comparison. The first hash is written by ``bootstrap_secrets`` at
|
||||
startup (generated, or migrated from a legacy password); here it is seeded
|
||||
directly into the store, then login checks against it and a password change
|
||||
re-hashes so the old password stops working and the new one starts.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from httpx import AsyncClient, Response
|
||||
|
||||
from routstr.core.db import create_session, set_admin_password
|
||||
|
||||
|
||||
async def _seed_password(password: str) -> None:
|
||||
# bootstrap_secrets owns first-password creation at startup; tests seed the
|
||||
# hash straight into the store the same way, then exercise login/change.
|
||||
async with create_session() as session:
|
||||
await set_admin_password(session, password)
|
||||
|
||||
|
||||
async def _login(client: AsyncClient, password: str) -> Response:
|
||||
return await client.post("/admin/api/login", json={"password": password})
|
||||
|
||||
|
||||
# --- login -----------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_login_500_when_no_password_configured(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
resp = await _login(integration_client, "anything")
|
||||
assert resp.status_code == 500
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_login_succeeds_with_correct_password(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
await _seed_password("correct horse")
|
||||
resp = await _login(integration_client, "correct horse")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["ok"] is True
|
||||
assert isinstance(body["token"], str) and body["token"]
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_login_rejects_wrong_password(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
await _seed_password("correct horse")
|
||||
resp = await _login(integration_client, "wrong horse")
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
# --- password change -------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_password_rehashes_so_only_new_works(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
await _seed_password("old password")
|
||||
login = await _login(integration_client, "old password")
|
||||
token = login.json()["token"]
|
||||
integration_client.headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
resp = await integration_client.patch(
|
||||
"/admin/api/password",
|
||||
json={"current_password": "old password", "new_password": "new password"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
|
||||
# Drop admin auth so the login calls aren't treated as authenticated noise.
|
||||
integration_client.headers.pop("Authorization", None)
|
||||
assert (await _login(integration_client, "old password")).status_code == 401
|
||||
assert (await _login(integration_client, "new password")).status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_password_rejects_wrong_current(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
await _seed_password("old password")
|
||||
login = await _login(integration_client, "old password")
|
||||
token = login.json()["token"]
|
||||
integration_client.headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
resp = await integration_client.patch(
|
||||
"/admin/api/password",
|
||||
json={"current_password": "not the password", "new_password": "new password"},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_password_rejects_short_new(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
await _seed_password("old password")
|
||||
login = await _login(integration_client, "old password")
|
||||
token = login.json()["token"]
|
||||
integration_client.headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
resp = await integration_client.patch(
|
||||
"/admin/api/password",
|
||||
json={"current_password": "old password", "new_password": "x"},
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
@@ -0,0 +1,113 @@
|
||||
"""Tests for the admin nsec rotation endpoint (issue #553).
|
||||
|
||||
The Nostr identity is a secret: it lives encrypted in the Secret store, never in
|
||||
the settings blob, so it cannot be set through the general settings PATCH. This
|
||||
dedicated endpoint is the supported way to set/rotate/clear it — it encrypts the
|
||||
key at rest, updates the live runtime identity (so signing picks it up without a
|
||||
restart), and derives the npub. Invalid keys are rejected.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
import time
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from httpx import AsyncClient
|
||||
|
||||
from routstr.core import vault
|
||||
from routstr.core.admin import admin_sessions
|
||||
from routstr.core.db import AsyncSession, get_secret
|
||||
from routstr.core.settings import derive_npub_from_nsec, settings
|
||||
|
||||
# A valid 64-char hex private key (accepted by nsec_to_keypair, as in bootstrap).
|
||||
NSEC_HEX = "1" * 64
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def admin_client(
|
||||
integration_client: AsyncClient,
|
||||
) -> AsyncGenerator[AsyncClient, None]:
|
||||
"""An integration_client pre-authenticated with an admin session token."""
|
||||
token = secrets.token_urlsafe(24)
|
||||
admin_sessions[token] = int(time.time()) + 3600
|
||||
integration_client.headers["Authorization"] = f"Bearer {token}"
|
||||
yield integration_client
|
||||
admin_sessions.pop(token, None)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_nsec_stores_encrypted_and_derives_npub(
|
||||
admin_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "nsec", "")
|
||||
monkeypatch.setattr(settings, "npub", "")
|
||||
|
||||
resp = await admin_client.patch("/admin/api/nsec", json={"nsec": NSEC_HEX})
|
||||
assert resp.status_code == 200
|
||||
|
||||
expected_npub = derive_npub_from_nsec(NSEC_HEX)
|
||||
assert resp.json() == {"ok": True, "npub": expected_npub}
|
||||
|
||||
# Stored encrypted at rest, decryptable back to the original key.
|
||||
integration_session.expunge_all()
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.encrypted_nsec is not None
|
||||
assert vault.is_encrypted(secret.encrypted_nsec)
|
||||
assert vault.decrypt(secret.encrypted_nsec) == NSEC_HEX
|
||||
|
||||
# Live runtime identity updated so Nostr signing reflects it without restart.
|
||||
assert settings.nsec == NSEC_HEX
|
||||
assert settings.npub == expected_npub
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_nsec_rejects_invalid_key(
|
||||
admin_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "nsec", "")
|
||||
monkeypatch.setattr(settings, "npub", "")
|
||||
|
||||
resp = await admin_client.patch(
|
||||
"/admin/api/nsec", json={"nsec": "not-a-real-nsec"}
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
|
||||
# Nothing stored, live identity untouched.
|
||||
integration_session.expunge_all()
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.encrypted_nsec is None
|
||||
assert settings.nsec == ""
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_nsec_clears_identity_with_empty_value(
|
||||
admin_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# Start from a node that has an identity...
|
||||
monkeypatch.setattr(settings, "nsec", "")
|
||||
monkeypatch.setattr(settings, "npub", "")
|
||||
set_resp = await admin_client.patch("/admin/api/nsec", json={"nsec": NSEC_HEX})
|
||||
assert set_resp.status_code == 200
|
||||
|
||||
# ...then clear it.
|
||||
clear_resp = await admin_client.patch("/admin/api/nsec", json={"nsec": ""})
|
||||
assert clear_resp.status_code == 200
|
||||
assert clear_resp.json() == {"ok": True, "npub": ""}
|
||||
|
||||
integration_session.expunge_all()
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.encrypted_nsec is None
|
||||
assert settings.nsec == ""
|
||||
assert settings.npub == ""
|
||||
@@ -0,0 +1,82 @@
|
||||
"""Tests for the admin settings endpoint's handling of secrets (issue #553).
|
||||
|
||||
``admin_password`` is no longer a settings field (it lives only as a one-way
|
||||
hash in the Secret store), so it must never appear in the GET/PATCH payloads.
|
||||
``nsec`` (in-memory at runtime) and ``upstream_api_key`` (still in the settings
|
||||
blob) are both redacted on read and ignored on write — they cannot be set
|
||||
through the general settings endpoint, only through their dedicated paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
import time
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from httpx import AsyncClient
|
||||
|
||||
from routstr.core.admin import admin_sessions
|
||||
from routstr.core.db import AsyncSession
|
||||
from routstr.core.settings import SettingsService, settings
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def admin_client(
|
||||
integration_client: AsyncClient,
|
||||
) -> AsyncGenerator[AsyncClient, None]:
|
||||
"""An integration_client pre-authenticated with an admin session token."""
|
||||
token = secrets.token_urlsafe(24)
|
||||
admin_sessions[token] = int(time.time()) + 3600
|
||||
integration_client.headers["Authorization"] = f"Bearer {token}"
|
||||
yield integration_client
|
||||
admin_sessions.pop(token, None)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_settings_omits_admin_password_and_redacts_secrets(
|
||||
admin_client: AsyncClient, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "nsec", "nsec-secret")
|
||||
monkeypatch.setattr(settings, "upstream_api_key", "sk-secret")
|
||||
|
||||
resp = await admin_client.get("/admin/api/settings")
|
||||
assert resp.status_code == 200
|
||||
|
||||
data = resp.json()
|
||||
assert "admin_password" not in data
|
||||
assert data["nsec"] == "[REDACTED]"
|
||||
assert data["upstream_api_key"] == "[REDACTED]"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_settings_ignores_secret_fields(
|
||||
admin_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# The PATCH path persists through SettingsService, which needs an
|
||||
# initialized current snapshot and a settings row in the shared test DB.
|
||||
await SettingsService.initialize(integration_session)
|
||||
monkeypatch.setattr(settings, "nsec", "original-nsec")
|
||||
|
||||
resp = await admin_client.patch(
|
||||
"/admin/api/settings",
|
||||
json={
|
||||
"name": "Renamed",
|
||||
"nsec": "attacker-nsec",
|
||||
"upstream_api_key": "attacker-key",
|
||||
"admin_password": "attacker-pw",
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
data = resp.json()
|
||||
assert data["name"] == "Renamed"
|
||||
assert "admin_password" not in data
|
||||
assert data["nsec"] == "[REDACTED]"
|
||||
# The live secret was not overwritten through the general settings endpoint.
|
||||
assert settings.nsec == "original-nsec"
|
||||
@@ -82,7 +82,9 @@ async def test_balance_never_negative_when_cost_exceeds_reservation(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
|
||||
await adjust_payment_for_tokens(
|
||||
key, response_data, integration_session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await _refresh(integration_session, key)
|
||||
|
||||
@@ -116,7 +118,9 @@ async def test_balance_floor_at_zero_on_overrun(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
|
||||
await adjust_payment_for_tokens(
|
||||
key, response_data, integration_session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await _refresh(integration_session, key)
|
||||
|
||||
@@ -155,7 +159,9 @@ async def test_full_cost_charged_when_balance_sufficient_for_overrun(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
|
||||
await adjust_payment_for_tokens(
|
||||
key, response_data, integration_session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await _refresh(integration_session, key)
|
||||
|
||||
@@ -223,15 +229,18 @@ async def test_concurrent_cost_overruns_never_negative(
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(ApiKey, key_hash)
|
||||
assert fresh_key is not None
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, response_data, session, deducted_max_cost
|
||||
)
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, response_data, session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await asyncio.gather(*[finalize() for _ in range(n_requests)])
|
||||
# Patch once around the gather: entering the same patch target from
|
||||
# concurrent tasks un-patches in the wrong order and leaks the mock into
|
||||
# every later test in the session.
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await asyncio.gather(*[finalize() for _ in range(n_requests)])
|
||||
|
||||
async with create_session() as session:
|
||||
final_key = await session.get(ApiKey, key_hash)
|
||||
@@ -279,7 +288,9 @@ async def test_zero_free_balance_overrun_is_safe(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
|
||||
await adjust_payment_for_tokens(
|
||||
key, response_data, integration_session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await _refresh(integration_session, key)
|
||||
|
||||
@@ -344,15 +355,18 @@ async def test_parallel_requests_no_free_inference(
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(ApiKey, key_hash)
|
||||
assert fresh_key is not None
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, response_data, session, deducted_max_cost
|
||||
)
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, response_data, session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await asyncio.gather(finalize(), finalize())
|
||||
# Patch once around the gather: entering the same patch target from two
|
||||
# concurrent tasks un-patches in the wrong order and leaks the mock into
|
||||
# every later test in the session.
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await asyncio.gather(finalize(), finalize())
|
||||
|
||||
async with create_session() as session:
|
||||
final_key = await session.get(ApiKey, key_hash)
|
||||
|
||||
@@ -77,7 +77,7 @@ async def test_child_key_flow(integration_session: AsyncSession) -> None:
|
||||
|
||||
try:
|
||||
adjustment = await adjust_payment_for_tokens(
|
||||
child_key_db, response_data, integration_session, 500
|
||||
child_key_db, response_data, integration_session, 500, None, None
|
||||
)
|
||||
assert adjustment["total_msats"] == 400
|
||||
|
||||
|
||||
@@ -0,0 +1,546 @@
|
||||
"""Failover requests are billed and forwarded as the provider that served them.
|
||||
|
||||
Covers the whole-system settlement path when two enabled providers expose the
|
||||
same model under different spellings and prices: the routing winner fails with
|
||||
a 502, the fallback provider serves, and the response must be billed at the
|
||||
fallback's configured rate, carry the fallback's model id in the forwarded
|
||||
request body, and echo the fallback's model id to the client.
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, AsyncGenerator
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from httpx import AsyncClient
|
||||
|
||||
from routstr.payment.models import Architecture, Model, Pricing
|
||||
from routstr.proxy import refresh_model_maps
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
CHEAP_BASE_URL = "https://cheap.example.com/v1"
|
||||
EXPENSIVE_BASE_URL = "https://expensive.example.com/v1"
|
||||
|
||||
|
||||
def _make_model(
|
||||
model_id: str,
|
||||
prompt_sats: float,
|
||||
completion_sats: float,
|
||||
max_cost: float = 50.0,
|
||||
) -> Model:
|
||||
"""Build a model whose USD and sats pricing rank consistently."""
|
||||
return Model(
|
||||
id=model_id,
|
||||
name=model_id,
|
||||
created=1,
|
||||
description="test model",
|
||||
context_length=8192,
|
||||
architecture=Architecture(
|
||||
modality="text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="gpt",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(
|
||||
prompt=prompt_sats, completion=completion_sats, max_cost=max_cost
|
||||
),
|
||||
sats_pricing=Pricing(
|
||||
prompt=prompt_sats, completion=completion_sats, max_cost=max_cost
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class _StaticProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider with a fixed model catalog and no remote refresh."""
|
||||
|
||||
def __init__(
|
||||
self, base_url: str, api_key: str, fee: float, model: Model
|
||||
) -> None:
|
||||
super().__init__(base_url, api_key, fee)
|
||||
self.provider_type = "custom"
|
||||
self._static_model = model
|
||||
|
||||
def get_cached_models(self) -> list[Model]:
|
||||
return [self._static_model]
|
||||
|
||||
async def refresh_models_cache(self) -> None:
|
||||
pass
|
||||
|
||||
|
||||
async def _install_providers(
|
||||
providers: list[_StaticProvider],
|
||||
) -> AsyncGenerator[None, None]:
|
||||
"""Install providers into the routing maps, restoring the originals after."""
|
||||
from routstr import proxy
|
||||
|
||||
original_upstreams = proxy.get_upstreams()
|
||||
with patch("routstr.proxy._upstreams", providers):
|
||||
await refresh_model_maps()
|
||||
yield
|
||||
with patch("routstr.proxy._upstreams", original_upstreams):
|
||||
await refresh_model_maps()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def dual_provider_maps(
|
||||
patched_db_engine: None,
|
||||
) -> AsyncGenerator[tuple[_StaticProvider, _StaticProvider], None]:
|
||||
"""Two same-tail providers under different spellings and prices."""
|
||||
cheap = _StaticProvider(
|
||||
CHEAP_BASE_URL,
|
||||
"key-cheap",
|
||||
1.0,
|
||||
_make_model("prova/dual-model", 0.001, 0.002),
|
||||
)
|
||||
expensive = _StaticProvider(
|
||||
EXPENSIVE_BASE_URL,
|
||||
"key-expensive",
|
||||
1.0,
|
||||
_make_model("provb/dual-model", 0.005, 0.010),
|
||||
)
|
||||
async for _ in _install_providers([cheap, expensive]):
|
||||
yield cheap, expensive
|
||||
|
||||
|
||||
def _upstream_response(request: httpx.Request) -> httpx.Response:
|
||||
"""502 from the cheap (winning) provider; a served completion elsewhere."""
|
||||
if request.url.host == "cheap.example.com":
|
||||
return httpx.Response(
|
||||
502,
|
||||
content=json.dumps({"error": {"message": "bad gateway"}}).encode(),
|
||||
headers={"content-type": "application/json"},
|
||||
)
|
||||
body = {
|
||||
"id": "chatcmpl-served",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "dual-model",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "hi"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 1000,
|
||||
"completion_tokens": 500,
|
||||
"total_tokens": 1500,
|
||||
},
|
||||
}
|
||||
return httpx.Response(
|
||||
200,
|
||||
content=json.dumps(body).encode(),
|
||||
headers={"content-type": "application/json"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_failover_serve_billed_at_serving_providers_rate(
|
||||
authenticated_client: AsyncClient,
|
||||
dual_provider_maps: tuple[_StaticProvider, _StaticProvider],
|
||||
) -> None:
|
||||
"""A fallback serve is billed at the fallback's price, not the winner's.
|
||||
|
||||
The cheap provider ranks first for the shared tail; it 502s and the
|
||||
expensive provider serves 1000 input + 500 output tokens. At the serving
|
||||
provider's sats pricing (0.005/0.010 sats per token) that is 10_000 msats;
|
||||
at the winner's (0.001/0.002) it would be 2_000 msats.
|
||||
"""
|
||||
sent_requests: list[httpx.Request] = []
|
||||
|
||||
# Patch the network transport (not AsyncClient.send) so the in-process
|
||||
# ASGI test client is untouched and only the proxy's upstream hop is mocked.
|
||||
async def fake_transport(
|
||||
request: httpx.Request, *args: Any, **kwargs: Any
|
||||
) -> httpx.Response:
|
||||
sent_requests.append(request)
|
||||
return _upstream_response(request)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
side_effect=fake_transport,
|
||||
),
|
||||
# cost_calculation binds sats_usd_price at import time, so the price
|
||||
# patch in the app fixture does not reach it; patch its own binding.
|
||||
patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price",
|
||||
return_value=0.0005,
|
||||
),
|
||||
):
|
||||
response = await authenticated_client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": "dual-model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = response.json()
|
||||
|
||||
# Both providers were attempted, cheapest first.
|
||||
assert [r.url.host for r in sent_requests] == [
|
||||
"cheap.example.com",
|
||||
"expensive.example.com",
|
||||
]
|
||||
|
||||
# The fallback must be asked for ITS OWN model spelling, not the winner's.
|
||||
forwarded_body = json.loads(sent_requests[1].content)
|
||||
assert forwarded_body["model"] == "provb/dual-model"
|
||||
|
||||
# The response echo names the model that actually served.
|
||||
assert payload["model"] == "provb/dual-model"
|
||||
|
||||
# Billed at the serving provider's rate: 1000/1000*5000 + 500/1000*10000.
|
||||
assert payload["cost"]["total_msats"] == 10_000
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def same_id_provider_maps(
|
||||
patched_db_engine: None,
|
||||
) -> AsyncGenerator[None, None]:
|
||||
"""Two providers exposing the IDENTICAL model id at different prices."""
|
||||
cheap = _StaticProvider(
|
||||
CHEAP_BASE_URL,
|
||||
"key-cheap",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.001, 0.002),
|
||||
)
|
||||
expensive = _StaticProvider(
|
||||
EXPENSIVE_BASE_URL,
|
||||
"key-expensive",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.005, 0.010),
|
||||
)
|
||||
async for _ in _install_providers([cheap, expensive]):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_same_id_failover_settles_at_serving_price(
|
||||
authenticated_client: AsyncClient,
|
||||
same_id_provider_maps: None,
|
||||
) -> None:
|
||||
"""Settlement must not re-derive pricing from the response's model string.
|
||||
|
||||
Both providers expose the exact same model id, so the forwarded body is
|
||||
identical either way — the only observable difference is the settled
|
||||
amount. The response's model string resolves to the alias winner (cheap),
|
||||
but the expensive provider served, so the bill must be 10_000 msats, not
|
||||
the winner's 2_000.
|
||||
"""
|
||||
sent_requests: list[httpx.Request] = []
|
||||
|
||||
async def fake_transport(
|
||||
request: httpx.Request, *args: Any, **kwargs: Any
|
||||
) -> httpx.Response:
|
||||
sent_requests.append(request)
|
||||
return _upstream_response(request)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
side_effect=fake_transport,
|
||||
),
|
||||
patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price",
|
||||
return_value=0.0005,
|
||||
),
|
||||
):
|
||||
response = await authenticated_client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": "dual-model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [r.url.host for r in sent_requests] == [
|
||||
"cheap.example.com",
|
||||
"expensive.example.com",
|
||||
]
|
||||
assert response.json()["cost"]["total_msats"] == 10_000
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_version_suffixed_model_id_routes(
|
||||
authenticated_client: AsyncClient,
|
||||
same_id_provider_maps: None,
|
||||
) -> None:
|
||||
"""A version-suffixed request (``…-YYYYMMDD``) routes to the base model.
|
||||
|
||||
Model resolution stripped the suffix but the provider lookup did not, so
|
||||
such requests resolved a model yet found no provider and 400'd. With the
|
||||
unified candidate lookup the strip applies to both.
|
||||
"""
|
||||
|
||||
async def fake_transport(
|
||||
request: httpx.Request, *args: Any, **kwargs: Any
|
||||
) -> httpx.Response:
|
||||
return _upstream_response(request)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
side_effect=fake_transport,
|
||||
),
|
||||
patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price",
|
||||
return_value=0.0005,
|
||||
),
|
||||
):
|
||||
response = await authenticated_client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": "dual-model-20260101",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def fee_split_provider_maps(
|
||||
patched_db_engine: None,
|
||||
) -> AsyncGenerator[None, None]:
|
||||
"""Same-tail providers whose fees differ; the serving one charges 1.5x."""
|
||||
cheap = _StaticProvider(
|
||||
CHEAP_BASE_URL,
|
||||
"key-cheap",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.001, 0.002),
|
||||
)
|
||||
expensive = _StaticProvider(
|
||||
EXPENSIVE_BASE_URL,
|
||||
"key-expensive",
|
||||
1.5,
|
||||
_make_model("dual-model", 0.005, 0.010),
|
||||
)
|
||||
async for _ in _install_providers([cheap, expensive]):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_usd_cost_serve_carries_serving_providers_fee(
|
||||
authenticated_client: AsyncClient,
|
||||
fee_split_provider_maps: None,
|
||||
) -> None:
|
||||
"""The USD-cost billing path applies the SERVING provider's fee.
|
||||
|
||||
The upstream that serves reports ``usage.cost`` in USD, so billing goes
|
||||
through the USD-cost path where the provider fee is applied explicitly.
|
||||
The serving provider's fee is 1.5; the alias winner's is 1.0. At 0.001 USD
|
||||
reported cost and 0.0005 USD/sat: 0.001 * 1.5 / 0.0005 = 3 sats = 3000
|
||||
msats (fee 1.0 would give 2000).
|
||||
"""
|
||||
sent_requests: list[httpx.Request] = []
|
||||
|
||||
def usd_cost_response(request: httpx.Request) -> httpx.Response:
|
||||
if request.url.host == "cheap.example.com":
|
||||
return httpx.Response(
|
||||
502,
|
||||
content=json.dumps(
|
||||
{"error": {"message": "bad gateway"}}
|
||||
).encode(),
|
||||
headers={"content-type": "application/json"},
|
||||
)
|
||||
body = {
|
||||
"id": "chatcmpl-usd",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "dual-model",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "hi"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 50,
|
||||
"total_tokens": 150,
|
||||
"cost": 0.001,
|
||||
},
|
||||
}
|
||||
return httpx.Response(
|
||||
200,
|
||||
content=json.dumps(body).encode(),
|
||||
headers={"content-type": "application/json"},
|
||||
)
|
||||
|
||||
async def fake_transport(
|
||||
request: httpx.Request, *args: Any, **kwargs: Any
|
||||
) -> httpx.Response:
|
||||
sent_requests.append(request)
|
||||
return usd_cost_response(request)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
side_effect=fake_transport,
|
||||
),
|
||||
patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price",
|
||||
return_value=0.0005,
|
||||
),
|
||||
):
|
||||
response = await authenticated_client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": "dual-model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [r.url.host for r in sent_requests] == [
|
||||
"cheap.example.com",
|
||||
"expensive.example.com",
|
||||
]
|
||||
assert response.json()["cost"]["total_msats"] == 3_000
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def envelope_split_provider_maps(
|
||||
patched_db_engine: None,
|
||||
) -> AsyncGenerator[None, None]:
|
||||
"""Same-id providers where the fallback's max cost dwarfs the key balance."""
|
||||
cheap = _StaticProvider(
|
||||
CHEAP_BASE_URL,
|
||||
"key-cheap",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
|
||||
)
|
||||
expensive = _StaticProvider(
|
||||
EXPENSIVE_BASE_URL,
|
||||
"key-expensive",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.005, 0.010, max_cost=20_000.0),
|
||||
)
|
||||
async for _ in _install_providers([cheap, expensive]):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_failover_beyond_balance_envelope_is_rejected(
|
||||
authenticated_client: AsyncClient,
|
||||
envelope_split_provider_maps: None,
|
||||
) -> None:
|
||||
"""A fallback whose max-cost envelope exceeds the balance is not served.
|
||||
|
||||
Admission and reservation are sized to the best-ranked candidate's max
|
||||
cost. When that candidate fails and the next one's envelope exceeds the
|
||||
key's balance, serving it could settle far beyond what admission allowed,
|
||||
so the request must be rejected (as it would be if the pricier candidate
|
||||
were ranked first) instead of forwarded.
|
||||
"""
|
||||
sent_requests: list[httpx.Request] = []
|
||||
|
||||
async def fake_transport(
|
||||
request: httpx.Request, *args: Any, **kwargs: Any
|
||||
) -> httpx.Response:
|
||||
sent_requests.append(request)
|
||||
return _upstream_response(request)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
side_effect=fake_transport,
|
||||
),
|
||||
patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price",
|
||||
return_value=0.0005,
|
||||
),
|
||||
):
|
||||
response = await authenticated_client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": "dual-model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
# The 20_000-sat envelope exceeds the key's 10_000-sat balance: the
|
||||
# fallback must be rejected before its upstream is ever contacted.
|
||||
assert response.status_code == 402
|
||||
assert [r.url.host for r in sent_requests] == ["cheap.example.com"]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def raised_envelope_provider_maps(
|
||||
patched_db_engine: None,
|
||||
) -> AsyncGenerator[None, None]:
|
||||
"""Same-id providers where the fallback needs a larger, affordable reserve."""
|
||||
cheap = _StaticProvider(
|
||||
CHEAP_BASE_URL,
|
||||
"key-cheap",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
|
||||
)
|
||||
expensive = _StaticProvider(
|
||||
EXPENSIVE_BASE_URL,
|
||||
"key-expensive",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.005, 0.010, max_cost=100.0),
|
||||
)
|
||||
async for _ in _install_providers([cheap, expensive]):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_failover_reserves_serving_candidates_envelope(
|
||||
authenticated_client: AsyncClient,
|
||||
raised_envelope_provider_maps: None,
|
||||
) -> None:
|
||||
"""An affordable pricier fallback is re-reserved, served, and billed.
|
||||
|
||||
The fallback's max cost (100 sats) exceeds the winner's (50 sats) but fits
|
||||
the key's balance, so the reservation is raised to the serving candidate's
|
||||
envelope and the request completes, billed at the serving rate with the
|
||||
unused reserve refunded.
|
||||
"""
|
||||
sent_requests: list[httpx.Request] = []
|
||||
|
||||
async def fake_transport(
|
||||
request: httpx.Request, *args: Any, **kwargs: Any
|
||||
) -> httpx.Response:
|
||||
sent_requests.append(request)
|
||||
return _upstream_response(request)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
side_effect=fake_transport,
|
||||
),
|
||||
patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price",
|
||||
return_value=0.0005,
|
||||
),
|
||||
):
|
||||
response = await authenticated_client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": "dual-model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [r.url.host for r in sent_requests] == [
|
||||
"cheap.example.com",
|
||||
"expensive.example.com",
|
||||
]
|
||||
assert response.json()["cost"]["total_msats"] == 10_000
|
||||
@@ -58,7 +58,7 @@ async def test_overrun_charges_after_reservation_swept(
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(
|
||||
key, response_data, integration_session, deducted_max_cost
|
||||
key, response_data, integration_session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
await integration_session.refresh(key)
|
||||
@@ -129,7 +129,7 @@ async def test_free_response_path_closed_end_to_end(
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(
|
||||
key, response_data, session, deducted_max_cost
|
||||
key, response_data, session, deducted_max_cost, None, None
|
||||
)
|
||||
|
||||
async with create_session() as session:
|
||||
|
||||
@@ -241,8 +241,10 @@ async def test_http_402_response_shape_on_insufficient_balance(
|
||||
mock_upstream.prepare_headers = MagicMock(return_value={})
|
||||
|
||||
with (
|
||||
patch("routstr.proxy.get_model_instance", return_value=mock_model),
|
||||
patch("routstr.proxy.get_provider_for_model", return_value=[mock_upstream]),
|
||||
patch(
|
||||
"routstr.proxy.get_candidates",
|
||||
return_value=[(mock_model, mock_upstream)],
|
||||
),
|
||||
# Patch where it is used (proxy imports it at module level)
|
||||
patch(
|
||||
"routstr.proxy.get_max_cost_for_model",
|
||||
|
||||
@@ -120,7 +120,9 @@ async def test_finalise_releases_reservation_and_charges_balance(
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
|
||||
|
||||
with patch("routstr.auth.calculate_cost", return_value=cost_data):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, cost)
|
||||
await adjust_payment_for_tokens(
|
||||
key, response_data, integration_session, cost, None, None
|
||||
)
|
||||
|
||||
await integration_session.refresh(key)
|
||||
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
"""Tests for the ``reset_admin_password`` recovery script (issue #553).
|
||||
|
||||
The script is the lockout escape hatch: it works without ``ROUTSTR_SECRET_KEY``
|
||||
(scrypt hashing is key-independent). Two explicit, mutually exclusive actions —
|
||||
``--password`` sets a new hash now, ``--regenerate`` clears the hash so the next
|
||||
boot generates and logs a fresh one. A bare invocation is informational only and
|
||||
must never touch the database (so nobody resets their password by accident).
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core import vault
|
||||
from routstr.core.db import get_secret, set_admin_password
|
||||
from scripts.reset_admin_password import apply_reset, build_parser, main
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_sets_a_verifiable_hash(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
await apply_reset(integration_session, password="recover-me-123")
|
||||
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.admin_password_hash is not None
|
||||
assert vault.verify_password("recover-me-123", secret.admin_password_hash) is True
|
||||
assert secret.updated_at is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_regenerate_clears_the_hash(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
# Start from a node that already has an admin password set.
|
||||
await set_admin_password(integration_session, "old-password-9")
|
||||
assert (await get_secret(integration_session)).admin_password_hash is not None
|
||||
|
||||
await apply_reset(integration_session, regenerate=True)
|
||||
|
||||
secret = await get_secret(integration_session)
|
||||
# Cleared -> the next boot's bootstrap_secrets generates and logs a new one.
|
||||
assert secret.admin_password_hash is None
|
||||
assert secret.updated_at is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_password_below_min_length_is_rejected(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
await set_admin_password(integration_session, "old-password-9")
|
||||
|
||||
with pytest.raises(ValueError, match="8 characters"):
|
||||
await apply_reset(integration_session, password="short")
|
||||
|
||||
# The existing password is untouched by the rejected reset.
|
||||
secret = await get_secret(integration_session)
|
||||
assert vault.verify_password("old-password-9", secret.admin_password_hash or "")
|
||||
|
||||
|
||||
def test_password_and_regenerate_are_mutually_exclusive() -> None:
|
||||
parser = build_parser()
|
||||
with pytest.raises(SystemExit):
|
||||
parser.parse_args(["--password", "abcd1234", "--regenerate"])
|
||||
|
||||
|
||||
def test_no_args_prints_help_and_never_opens_a_session(
|
||||
capsys: pytest.CaptureFixture[str],
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
def _fail() -> None:
|
||||
raise AssertionError("a bare invocation must not touch the database")
|
||||
|
||||
monkeypatch.setattr("scripts.reset_admin_password.create_session", _fail)
|
||||
|
||||
assert main([]) == 0
|
||||
assert "usage" in capsys.readouterr().out.lower()
|
||||
@@ -0,0 +1,484 @@
|
||||
"""Tests for ``bootstrap_secrets`` — moving node secrets into the Secret store.
|
||||
|
||||
Specifies the per-secret bootstrap that runs at startup (issue #553). For both
|
||||
the admin password and the nsec it follows the same three branches: use the
|
||||
column if already set, otherwise migrate any legacy plaintext (env first, then
|
||||
the old settings blob), otherwise — admin password only — generate and log one.
|
||||
A column written under a different ROUTSTR_SECRET_KEY fails fast rather than
|
||||
silently corrupting state.
|
||||
"""
|
||||
|
||||
import json
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from typing import Any, AsyncGenerator
|
||||
|
||||
import pytest
|
||||
from sqlmodel import text
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core import vault
|
||||
from routstr.core.db import NsecState, get_secret, set_nsec
|
||||
from routstr.core.settings import (
|
||||
SettingsService,
|
||||
bootstrap_secrets,
|
||||
derive_npub_from_nsec,
|
||||
settings,
|
||||
)
|
||||
|
||||
# Valid Fernet keys; must match the suite default in tests/conftest.py.
|
||||
TEST_SECRET_KEY = "l_Tkp-7xmjcQ-IFhr6qhILrU8HPRbEmYMrfSbo_5srU="
|
||||
TEST_SECRET_KEY_ALT = "_Teyrky_iToeDK51Tj1FsI9MJ340_cqKGmeher-a7MQ="
|
||||
|
||||
NSEC_HEX = "1" * 64
|
||||
# A different key, standing in for a stale value left behind in env/blob after
|
||||
# the vault has taken ownership of the real one.
|
||||
STALE_NSEC_HEX = "2" * 64
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def clean_secret_env(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""No ambient legacy secrets, and a known in-memory settings baseline."""
|
||||
monkeypatch.delenv("ADMIN_PASSWORD", raising=False)
|
||||
monkeypatch.delenv("NSEC", raising=False)
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY)
|
||||
monkeypatch.setattr(settings, "nsec", "")
|
||||
monkeypatch.setattr(settings, "npub", "")
|
||||
monkeypatch.setattr(settings, "http_url", "")
|
||||
|
||||
|
||||
async def _create_settings_blob(session: AsyncSession, data: dict) -> None:
|
||||
await session.exec( # type: ignore
|
||||
text(
|
||||
"CREATE TABLE IF NOT EXISTS settings "
|
||||
"(id INTEGER PRIMARY KEY, data TEXT NOT NULL, "
|
||||
"updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP)"
|
||||
)
|
||||
)
|
||||
await session.exec( # type: ignore
|
||||
text("INSERT INTO settings (id, data) VALUES (1, :data)").bindparams(
|
||||
data=json.dumps(data)
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
# --- admin password --------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generates_admin_password_when_none(
|
||||
clean_secret_env: None, integration_session: AsyncSession
|
||||
) -> None:
|
||||
await bootstrap_secrets(integration_session)
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.admin_password_hash is not None
|
||||
assert secret.admin_password_hash.startswith("scrypt:")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_password_generation_is_idempotent(
|
||||
clean_secret_env: None, integration_session: AsyncSession
|
||||
) -> None:
|
||||
await bootstrap_secrets(integration_session)
|
||||
first = (await get_secret(integration_session)).admin_password_hash
|
||||
await bootstrap_secrets(integration_session)
|
||||
second = (await get_secret(integration_session)).admin_password_hash
|
||||
assert first is not None and first == second
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hashes_legacy_admin_password_from_env(
|
||||
clean_secret_env: None,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("ADMIN_PASSWORD", "hunter2")
|
||||
await bootstrap_secrets(integration_session)
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.admin_password_hash is not None
|
||||
assert vault.verify_password("hunter2", secret.admin_password_hash) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_hashes_legacy_admin_password_from_blob(
|
||||
clean_secret_env: None, integration_session: AsyncSession
|
||||
) -> None:
|
||||
# No ADMIN_PASSWORD in env, but the old settings blob carries one.
|
||||
await _create_settings_blob(integration_session, {"admin_password": "blobpw"})
|
||||
await bootstrap_secrets(integration_session)
|
||||
secret = await get_secret(integration_session)
|
||||
assert vault.verify_password("blobpw", secret.admin_password_hash or "") is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_password_race_adopts_winner_without_clobber(
|
||||
clean_secret_env: None,
|
||||
integration_engine: Any,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
capsys: pytest.CaptureFixture[str],
|
||||
) -> None:
|
||||
# Two workers boot against one shared DB and both read a null admin password.
|
||||
# The first to commit "wins" and shows the operator its generated password. A
|
||||
# worker that read null but lost the race must NOT overwrite the winner's hash
|
||||
# (which the operator may already be logging in with) and must NOT print a
|
||||
# second password that will never work.
|
||||
#
|
||||
# The race window is forced deterministically: a hook fires inside bootstrap's
|
||||
# generate branch (so it only runs once this worker has committed to
|
||||
# generating) and commits the winner's password on a separate connection
|
||||
# before this worker writes its own.
|
||||
import sqlite3
|
||||
|
||||
from routstr.core import settings as settings_mod
|
||||
|
||||
db_file = integration_engine.url.database
|
||||
winner_hash = vault.hash_password("winner-password-123")
|
||||
real_token = settings_mod.secrets.token_urlsafe
|
||||
|
||||
def commit_winner_then_generate(nbytes: int) -> str:
|
||||
conn = sqlite3.connect(db_file)
|
||||
conn.execute(
|
||||
"UPDATE secrets SET admin_password_hash = ? WHERE id = 1", (winner_hash,)
|
||||
)
|
||||
conn.commit()
|
||||
conn.close()
|
||||
return real_token(nbytes)
|
||||
|
||||
monkeypatch.setattr(
|
||||
settings_mod.secrets, "token_urlsafe", commit_winner_then_generate
|
||||
)
|
||||
|
||||
await get_secret(integration_session) # row exists, password still null
|
||||
capsys.readouterr() # drop anything emitted before the race resolves
|
||||
await bootstrap_secrets(integration_session)
|
||||
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.admin_password_hash is not None
|
||||
# The winner's password survives and still verifies — no clobber.
|
||||
assert vault.verify_password("winner-password-123", secret.admin_password_hash)
|
||||
# The losing worker stayed silent — no second generated password was leaked.
|
||||
assert "generated a temporary" not in capsys.readouterr().out
|
||||
|
||||
|
||||
# --- nsec ------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_encrypts_legacy_nsec_from_env_and_derives_npub(
|
||||
clean_secret_env: None,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("NSEC", NSEC_HEX)
|
||||
await bootstrap_secrets(integration_session)
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.encrypted_nsec is not None
|
||||
assert vault.is_encrypted(secret.encrypted_nsec) is True
|
||||
assert vault.decrypt(secret.encrypted_nsec) == NSEC_HEX
|
||||
# In-memory runtime value is the decrypted nsec, and npub is derived from it.
|
||||
assert settings.nsec == NSEC_HEX
|
||||
assert settings.npub == derive_npub_from_nsec(NSEC_HEX)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decrypts_existing_nsec_column(
|
||||
clean_secret_env: None, integration_session: AsyncSession
|
||||
) -> None:
|
||||
secret = await get_secret(integration_session)
|
||||
secret.encrypted_nsec = vault.encrypt(NSEC_HEX)
|
||||
secret.nsec_state = NsecState.encrypted
|
||||
integration_session.add(secret)
|
||||
await integration_session.commit()
|
||||
stored = secret.encrypted_nsec
|
||||
|
||||
await bootstrap_secrets(integration_session)
|
||||
reloaded = await get_secret(integration_session)
|
||||
assert settings.nsec == NSEC_HEX
|
||||
# The column is reused, not re-encrypted.
|
||||
assert reloaded.encrypted_nsec == stored
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fail_fast_when_nsec_encrypted_with_different_key(
|
||||
clean_secret_env: None,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# Encrypt the column under the alternate key, then bootstrap under the
|
||||
# suite key -> the value cannot be decrypted -> clear startup failure.
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY_ALT)
|
||||
secret = await get_secret(integration_session)
|
||||
secret.encrypted_nsec = vault.encrypt(NSEC_HEX)
|
||||
secret.nsec_state = NsecState.encrypted
|
||||
integration_session.add(secret)
|
||||
await integration_session.commit()
|
||||
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY)
|
||||
with pytest.raises(RuntimeError, match="ROUTSTR_SECRET_KEY"):
|
||||
await bootstrap_secrets(integration_session)
|
||||
|
||||
|
||||
# --- encryption is mandatory, key custody is not: upgrade without a key --------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_legacy_nsec_without_secret_key_generates_and_encrypts(
|
||||
clean_secret_env: None,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path: Path,
|
||||
capsys: pytest.CaptureFixture[str],
|
||||
) -> None:
|
||||
# A node upgrading with a legacy plaintext NSEC but no ROUTSTR_SECRET_KEY must
|
||||
# NOT break. Encryption at rest stays mandatory (the nsec is never persisted
|
||||
# in plaintext), but the key custody is flexible: bootstrap generates a master
|
||||
# key, persists it to the key file, warns loudly, and encrypts the identity —
|
||||
# so the node keeps running instead of refusing to boot.
|
||||
monkeypatch.delenv("ROUTSTR_SECRET_KEY", raising=False)
|
||||
key_file = tmp_path / "routstr_secret.key"
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY_FILE", str(key_file))
|
||||
monkeypatch.setenv("NSEC", NSEC_HEX)
|
||||
|
||||
await bootstrap_secrets(integration_session)
|
||||
|
||||
# A master key was generated and persisted...
|
||||
assert key_file.exists()
|
||||
# ...the nsec is encrypted at rest under it, never stored in plaintext...
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.encrypted_nsec is not None
|
||||
assert vault.is_encrypted(secret.encrypted_nsec) is True
|
||||
assert vault.decrypt(secret.encrypted_nsec) == NSEC_HEX
|
||||
assert secret.nsec_state == NsecState.encrypted
|
||||
# ...the node holds the live identity (npub derived from it)...
|
||||
assert settings.nsec == NSEC_HEX
|
||||
assert settings.npub == derive_npub_from_nsec(NSEC_HEX)
|
||||
# ...and the operator is loudly told a key was generated and must be backed up
|
||||
# (path shown, but never the key value) so an upgrade cannot silently create
|
||||
# an unbacked key nor leak the key into captured stdout / aggregated logs.
|
||||
out = capsys.readouterr().out
|
||||
assert str(key_file) in out
|
||||
assert key_file.read_text().strip() not in out
|
||||
assert "BACK UP" in out.upper()
|
||||
|
||||
|
||||
# --- boot ordering: rescue legacy blob secrets before they are stripped ----
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blob_only_nsec_is_migrated_before_blob_is_stripped(
|
||||
clean_secret_env: None, integration_session: AsyncSession
|
||||
) -> None:
|
||||
# Legacy node whose nsec lives ONLY in the settings blob (never in env).
|
||||
# bootstrap_secrets must run *before* SettingsService.initialize strips the
|
||||
# blob, or the only copy of the secret would be lost.
|
||||
await _create_settings_blob(
|
||||
integration_session, {"nsec": NSEC_HEX, "name": "LegacyNode"}
|
||||
)
|
||||
|
||||
await bootstrap_secrets(integration_session)
|
||||
await SettingsService.initialize(integration_session)
|
||||
|
||||
secret = await get_secret(integration_session)
|
||||
# The plaintext nsec has been moved into the encrypted Secret store...
|
||||
assert secret.encrypted_nsec is not None
|
||||
assert vault.decrypt(secret.encrypted_nsec) == NSEC_HEX
|
||||
assert settings.nsec == NSEC_HEX
|
||||
# ...and stripped from the persisted settings blob.
|
||||
row = await integration_session.exec( # type: ignore
|
||||
text("SELECT data FROM settings WHERE id = 1")
|
||||
)
|
||||
blob = json.loads(row.first()[0])
|
||||
assert "nsec" not in blob
|
||||
assert blob["name"] == "LegacyNode"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_does_not_clobber_store_only_nsec(
|
||||
clean_secret_env: None, integration_session: AsyncSession
|
||||
) -> None:
|
||||
# Steady state after migration: the nsec lives ONLY in the encrypted Secret
|
||||
# store (NSEC removed from env, blob already stripped on a previous boot).
|
||||
# bootstrap decrypts it into memory; initialize then re-derives settings from
|
||||
# the secret-free blob and must NOT wipe the live nsec back to empty (or the
|
||||
# node would silently stop signing Nostr announcements).
|
||||
await _create_settings_blob(integration_session, {"name": "LegacyNode"})
|
||||
secret = await get_secret(integration_session)
|
||||
secret.encrypted_nsec = vault.encrypt(NSEC_HEX)
|
||||
secret.nsec_state = NsecState.encrypted
|
||||
integration_session.add(secret)
|
||||
await integration_session.commit()
|
||||
|
||||
await bootstrap_secrets(integration_session)
|
||||
assert settings.nsec == NSEC_HEX # bootstrap decrypted it into memory
|
||||
|
||||
await SettingsService.initialize(integration_session)
|
||||
# The live secret survives initialize even though no env/blob carries it...
|
||||
assert settings.nsec == NSEC_HEX
|
||||
# ...and is still never written back to the persisted blob.
|
||||
row = await integration_session.exec( # type: ignore
|
||||
text("SELECT data FROM settings WHERE id = 1")
|
||||
)
|
||||
assert "nsec" not in json.loads(row.first()[0])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_env_nsec_does_not_override_vault_nsec(
|
||||
clean_secret_env: None,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# The vault owns the nsec, but a stale NSEC (e.g. the operator rotated the
|
||||
# key in the UI yet left the old value in .env) is still in the environment.
|
||||
# bootstrap decrypts the store value; initialize must NOT let the stale env
|
||||
# value clobber it, or a restart silently reverts to the old identity.
|
||||
await set_nsec(integration_session, NSEC_HEX)
|
||||
|
||||
monkeypatch.setenv("NSEC", STALE_NSEC_HEX)
|
||||
await _create_settings_blob(integration_session, {"name": "LegacyNode"})
|
||||
|
||||
await bootstrap_secrets(integration_session)
|
||||
await SettingsService.initialize(integration_session)
|
||||
|
||||
# The vault value wins; the stale env value is ignored.
|
||||
assert settings.nsec == NSEC_HEX
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_env_nsec_does_not_split_npub_from_vault_nsec(
|
||||
clean_secret_env: None,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# As above, the vault owns the nsec while a stale NSEC lingers in env. The
|
||||
# private key correctly comes from the vault, but the npub must too: if
|
||||
# initialize derives the public key from the stale env nsec, the node ends up
|
||||
# with a private key from the vault and a public key from the old env value,
|
||||
# and anything reading settings.npub announces the wrong Nostr identity.
|
||||
expected_npub = derive_npub_from_nsec(NSEC_HEX)
|
||||
stale_npub = derive_npub_from_nsec(STALE_NSEC_HEX)
|
||||
assert expected_npub and stale_npub and expected_npub != stale_npub # guard
|
||||
|
||||
await set_nsec(integration_session, NSEC_HEX)
|
||||
|
||||
monkeypatch.setenv("NSEC", STALE_NSEC_HEX)
|
||||
await _create_settings_blob(integration_session, {"name": "LegacyNode"})
|
||||
|
||||
await bootstrap_secrets(integration_session)
|
||||
await SettingsService.initialize(integration_session)
|
||||
|
||||
assert settings.nsec == NSEC_HEX
|
||||
assert settings.npub == expected_npub
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cleared_nsec_stays_cleared_across_reboot(
|
||||
clean_secret_env: None,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# An identity was imported from env, then the operator cleared it via the
|
||||
# admin API. The old NSEC is still in env. On the NEXT PROCESS the cleared
|
||||
# identity must stay cleared, not get resurrected from the stale env value.
|
||||
monkeypatch.setenv("NSEC", NSEC_HEX)
|
||||
await bootstrap_secrets(integration_session)
|
||||
assert settings.nsec == NSEC_HEX
|
||||
|
||||
# Clear via the admin path (store empty, vault owns it).
|
||||
await set_nsec(integration_session, "")
|
||||
|
||||
# Simulate a fresh process rather than pre-clearing the live singleton: the
|
||||
# pydantic settings global reloads the (still-stale) NSEC from env and derives
|
||||
# its npub, which is exactly the in-memory state a new boot starts from before
|
||||
# bootstrap runs. The cleared store must win over this stale live value.
|
||||
monkeypatch.setattr(settings, "nsec", NSEC_HEX)
|
||||
monkeypatch.setattr(settings, "npub", derive_npub_from_nsec(NSEC_HEX))
|
||||
|
||||
await bootstrap_secrets(integration_session)
|
||||
|
||||
reloaded = await get_secret(integration_session)
|
||||
assert reloaded.nsec_state == NsecState.cleared
|
||||
assert reloaded.encrypted_nsec is None # not re-imported
|
||||
assert settings.nsec == "" # stays cleared
|
||||
assert settings.npub == "" # and no derived public identity survives
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_keeps_npub_matching_store_only_nsec(
|
||||
clean_secret_env: None, integration_session: AsyncSession
|
||||
) -> None:
|
||||
# Steady state with mandatory encryption: the nsec lives ONLY in the
|
||||
# encrypted Secret store (env carries no NSEC) and the blob has no npub.
|
||||
# bootstrap decrypts the nsec and derives the npub into memory; initialize
|
||||
# then re-derives settings from the npub-less blob and must NOT wipe the npub
|
||||
# back to empty, or the node holds a private key with no matching public key
|
||||
# and silently stops announcing a usable Nostr identity.
|
||||
expected_npub = derive_npub_from_nsec(NSEC_HEX)
|
||||
assert expected_npub # guard: the test key must yield a real npub
|
||||
|
||||
await _create_settings_blob(integration_session, {"name": "LegacyNode"})
|
||||
secret = await get_secret(integration_session)
|
||||
secret.encrypted_nsec = vault.encrypt(NSEC_HEX)
|
||||
secret.nsec_state = NsecState.encrypted
|
||||
integration_session.add(secret)
|
||||
await integration_session.commit()
|
||||
|
||||
await bootstrap_secrets(integration_session)
|
||||
assert settings.npub == expected_npub # bootstrap derived it
|
||||
|
||||
await SettingsService.initialize(integration_session)
|
||||
# The npub still matches the live nsec...
|
||||
assert settings.nsec == NSEC_HEX
|
||||
assert settings.npub == expected_npub
|
||||
# ...and is persisted to the blob (it is public, not a stripped secret).
|
||||
row = await integration_session.exec( # type: ignore
|
||||
text("SELECT data FROM settings WHERE id = 1")
|
||||
)
|
||||
assert json.loads(row.first()[0])["npub"] == expected_npub
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_startup_runs_bootstrap_before_settings_initialize(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# The two tests above prove the migration outcome *given* the call order;
|
||||
# they hardcode that order themselves. This one guards the order at its real
|
||||
# call site — the application lifespan — so a reorder in main.py (which would
|
||||
# strip a blob-only secret before bootstrap could rescue it) is caught.
|
||||
import routstr.core.main as main
|
||||
|
||||
order: list[str] = []
|
||||
|
||||
class _Abort(Exception):
|
||||
pass
|
||||
|
||||
@asynccontextmanager
|
||||
async def fake_create_session() -> AsyncGenerator[None, None]:
|
||||
yield None
|
||||
|
||||
async def fake_bootstrap(session: Any) -> None:
|
||||
order.append("bootstrap")
|
||||
|
||||
async def fake_initialize(session: Any) -> None:
|
||||
order.append("initialize")
|
||||
# Stop startup here, before the background-task fan-out (prices, nostr,
|
||||
# upstreams) that we don't want to run in a unit test.
|
||||
raise _Abort()
|
||||
|
||||
async def noop_init_db() -> None:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(main, "configure_litellm", lambda: None)
|
||||
monkeypatch.setattr(main, "register_deepseek_v4_pricing", lambda: None)
|
||||
monkeypatch.setattr(main, "run_migrations", lambda: None)
|
||||
monkeypatch.setattr(main, "init_db", noop_init_db)
|
||||
monkeypatch.setattr(main, "create_session", fake_create_session)
|
||||
monkeypatch.setattr(main, "bootstrap_secrets", fake_bootstrap)
|
||||
monkeypatch.setattr(main.SettingsService, "initialize", fake_initialize)
|
||||
|
||||
with pytest.raises(_Abort):
|
||||
async with main.lifespan(main.app):
|
||||
pass
|
||||
|
||||
assert order == ["bootstrap", "initialize"]
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Tests for the ``Secret`` singleton model (issue #553).
|
||||
|
||||
Specifies the node-level secret store: a single row (``id=1``, like
|
||||
``RoutstrFee``) holding the one-way admin-password hash and the encrypted nsec.
|
||||
``get_secret`` is get-or-create, so callers always get the singleton without
|
||||
worrying whether it has been initialised yet. Encoding of the values themselves
|
||||
lives in ``routstr.core.vault``; here we only assert the row persists and stays
|
||||
a singleton.
|
||||
"""
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from sqlmodel import select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import Secret, get_secret
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_secret_creates_singleton(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
secret = await get_secret(integration_session)
|
||||
assert secret.id == 1
|
||||
# Fresh row carries no secret material yet.
|
||||
assert secret.admin_password_hash is None
|
||||
assert secret.encrypted_nsec is None
|
||||
assert secret.updated_at is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_secret_is_idempotent(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
first = await get_secret(integration_session)
|
||||
second = await get_secret(integration_session)
|
||||
assert first.id == second.id == 1
|
||||
rows = (await integration_session.exec(select(Secret))).all()
|
||||
assert len(rows) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_secret_fields_round_trip(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
secret = await get_secret(integration_session)
|
||||
secret.admin_password_hash = "scrypt:16384:8:1:c2FsdA==:aGFzaA=="
|
||||
secret.encrypted_nsec = "fernet:v1:gAAAAA"
|
||||
secret.updated_at = int(time.time())
|
||||
integration_session.add(secret)
|
||||
await integration_session.commit()
|
||||
|
||||
integration_session.expunge_all()
|
||||
reloaded = await get_secret(integration_session)
|
||||
assert reloaded.admin_password_hash == "scrypt:16384:8:1:c2FsdA==:aGFzaA=="
|
||||
assert reloaded.encrypted_nsec == "fernet:v1:gAAAAA"
|
||||
assert reloaded.updated_at is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_secret_tolerates_concurrent_first_insert(
|
||||
integration_engine: Any,
|
||||
integration_session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# A second worker wins the race and commits the singleton row first.
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as other:
|
||||
other.add(Secret(id=1, admin_password_hash="scrypt:from-other-worker"))
|
||||
await other.commit()
|
||||
|
||||
# Reproduce the race window: our session's first read still sees no row, so
|
||||
# it attempts to INSERT a duplicate id=1. The real IntegrityError that follows
|
||||
# must be recovered (roll back, re-read) rather than crashing startup.
|
||||
real_get = integration_session.get
|
||||
calls = {"n": 0}
|
||||
|
||||
async def stale_first_read(model: Any, pk: Any) -> Any:
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 1:
|
||||
return None
|
||||
return await real_get(model, pk)
|
||||
|
||||
monkeypatch.setattr(integration_session, "get", stale_first_read)
|
||||
|
||||
secret = await get_secret(integration_session)
|
||||
|
||||
# Recovered the other worker's row; no crash, still a single row.
|
||||
assert secret.id == 1
|
||||
assert secret.admin_password_hash == "scrypt:from-other-worker"
|
||||
rows = (await integration_session.exec(select(Secret))).all()
|
||||
assert len(rows) == 1
|
||||
@@ -0,0 +1,82 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.core import admin
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("requested_mint", [None, "https://secondary.example"])
|
||||
async def test_withdraw_uses_effective_mint_and_records_outgoing_transaction(
|
||||
monkeypatch: pytest.MonkeyPatch, requested_mint: str | None
|
||||
) -> None:
|
||||
primary_mint = "https://primary.example"
|
||||
effective_mint = requested_mint or primary_mint
|
||||
wallet = object()
|
||||
proofs = [SimpleNamespace(amount=40), SimpleNamespace(amount=60)]
|
||||
token = "cashuBoutgoing"
|
||||
|
||||
get_wallet = AsyncMock(return_value=wallet)
|
||||
get_proofs = Mock(return_value=proofs)
|
||||
filter_proofs = AsyncMock(return_value=proofs)
|
||||
send_token = AsyncMock(return_value=token)
|
||||
store_transaction = AsyncMock(return_value=True)
|
||||
|
||||
monkeypatch.setattr(admin, "get_wallet", get_wallet)
|
||||
monkeypatch.setattr(admin, "get_proofs_per_mint_and_unit", get_proofs)
|
||||
monkeypatch.setattr(admin, "slow_filter_spend_proofs", filter_proofs)
|
||||
monkeypatch.setattr(admin, "send_token", send_token)
|
||||
monkeypatch.setattr(admin, "store_cashu_transaction", store_transaction)
|
||||
monkeypatch.setattr(admin.settings, "primary_mint", primary_mint)
|
||||
|
||||
result = await admin.withdraw(
|
||||
Mock(),
|
||||
admin.WithdrawRequest(amount=75, mint_url=requested_mint, unit="sat"),
|
||||
)
|
||||
|
||||
assert result == {"token": token}
|
||||
get_wallet.assert_awaited_once_with(effective_mint, "sat")
|
||||
get_proofs.assert_called_once_with(wallet, effective_mint, "sat", not_reserved=True)
|
||||
filter_proofs.assert_awaited_once_with(proofs, wallet)
|
||||
send_token.assert_awaited_once_with(75, "sat", effective_mint)
|
||||
store_transaction.assert_awaited_once_with(
|
||||
token=token,
|
||||
amount=75,
|
||||
unit="sat",
|
||||
mint_url=effective_mint,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="admin",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_withdraw_returns_issued_token_when_audit_storage_fails(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
mint = "https://primary.example"
|
||||
proofs = [SimpleNamespace(amount=100)]
|
||||
token = "cashuBrecoverable"
|
||||
|
||||
monkeypatch.setattr(admin, "get_wallet", AsyncMock(return_value=object()))
|
||||
monkeypatch.setattr(
|
||||
admin, "get_proofs_per_mint_and_unit", Mock(return_value=proofs)
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
admin, "slow_filter_spend_proofs", AsyncMock(return_value=proofs)
|
||||
)
|
||||
monkeypatch.setattr(admin, "send_token", AsyncMock(return_value=token))
|
||||
monkeypatch.setattr(
|
||||
admin,
|
||||
"store_cashu_transaction",
|
||||
AsyncMock(side_effect=RuntimeError("database unavailable")),
|
||||
)
|
||||
critical = Mock()
|
||||
monkeypatch.setattr(admin.logger, "critical", critical)
|
||||
monkeypatch.setattr(admin.settings, "primary_mint", mint)
|
||||
|
||||
result = await admin.withdraw(Mock(), admin.WithdrawRequest(amount=75))
|
||||
|
||||
assert result == {"token": token}
|
||||
critical.assert_called_once()
|
||||
@@ -137,12 +137,12 @@ def test_create_model_mappings_includes_db_override_for_missing_cached_model(
|
||||
|
||||
model_instances, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[provider],
|
||||
overrides_by_id={"azure/gpt-4o": (override_row, 1.01)},
|
||||
disabled_model_ids=set(),
|
||||
overrides_by_key={("azure/gpt-4o", 7): (override_row, 1.01)},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert "azure/gpt-4o" in model_instances
|
||||
assert provider_map["azure/gpt-4o"] == [provider]
|
||||
assert [p for _, p in provider_map["azure/gpt-4o"]] == [provider]
|
||||
assert "gpt-4o" in unique_models
|
||||
|
||||
|
||||
@@ -182,11 +182,73 @@ def test_create_model_mappings_dedupes_with_provider_identity_not_provider_type(
|
||||
|
||||
_, provider_map, _ = create_model_mappings(
|
||||
upstreams=[provider_a, provider_b],
|
||||
overrides_by_id={"azure/gpt-4o": (override_row, 1.01)},
|
||||
disabled_model_ids=set(),
|
||||
overrides_by_key={("azure/gpt-4o", 2): (override_row, 1.01)},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
providers_for_alias = provider_map["azure/gpt-4o"]
|
||||
providers_for_alias = [p for _, p in provider_map["azure/gpt-4o"]]
|
||||
assert provider_a in providers_for_alias
|
||||
assert provider_b in providers_for_alias
|
||||
assert len(providers_for_alias) == 2
|
||||
|
||||
|
||||
def test_create_model_mappings_applies_override_only_to_matching_provider(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Same-id overrides must not add provider-specific aliases to other providers."""
|
||||
provider_a_model = create_test_model("same-id", prompt_price=0.01)
|
||||
provider_a = create_test_provider(
|
||||
"provider-a",
|
||||
"https://provider-a.example/v1",
|
||||
db_id=1,
|
||||
models=[provider_a_model],
|
||||
)
|
||||
provider_b_model = create_test_model("same-id", prompt_price=0.02)
|
||||
provider_b = create_test_provider(
|
||||
"provider-b",
|
||||
"https://provider-b.example/v1",
|
||||
db_id=2,
|
||||
models=[provider_b_model],
|
||||
)
|
||||
|
||||
override_model = create_test_model("same-id", prompt_price=0.001)
|
||||
override_model.alias_ids = ["provider-b-only"]
|
||||
override_row = SimpleNamespace(id="same-id", upstream_provider_id=2, enabled=True)
|
||||
|
||||
def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def]
|
||||
return override_model
|
||||
|
||||
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
|
||||
|
||||
_, provider_map, _ = create_model_mappings(
|
||||
upstreams=[provider_a, provider_b],
|
||||
overrides_by_key={("same-id", 2): (override_row, 1.01)},
|
||||
disabled_model_keys=set(),
|
||||
)
|
||||
|
||||
assert [p for _, p in provider_map["provider-b-only"]] == [provider_b]
|
||||
assert {p for _, p in provider_map["same-id"]} == {provider_a, provider_b}
|
||||
|
||||
|
||||
def test_create_model_mappings_disables_only_matching_provider() -> None:
|
||||
"""Disabled overrides are scoped to the provider row, not the shared model id."""
|
||||
provider_a = create_test_provider(
|
||||
"provider-a",
|
||||
"https://provider-a.example/v1",
|
||||
db_id=1,
|
||||
models=[create_test_model("same-id")],
|
||||
)
|
||||
provider_b = create_test_provider(
|
||||
"provider-b",
|
||||
"https://provider-b.example/v1",
|
||||
db_id=2,
|
||||
models=[create_test_model("same-id")],
|
||||
)
|
||||
|
||||
_, provider_map, _ = create_model_mappings(
|
||||
upstreams=[provider_a, provider_b],
|
||||
overrides_by_key={},
|
||||
disabled_model_keys={("same-id", 2)},
|
||||
)
|
||||
|
||||
assert [p for _, p in provider_map["same-id"]] == [provider_a]
|
||||
|
||||
@@ -0,0 +1,143 @@
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.core.db import CashuTransaction
|
||||
from routstr.upstream.auto_topup import _check_and_topup
|
||||
|
||||
|
||||
def _row() -> MagicMock:
|
||||
row = MagicMock()
|
||||
row.id = "provider-1"
|
||||
row.base_url = "https://provider.test"
|
||||
row.api_key = "secret"
|
||||
row.provider_settings = json.dumps(
|
||||
{
|
||||
"auto_topup": True,
|
||||
"topup_threshold": 100,
|
||||
"topup_amount_limit": 50,
|
||||
"topup_mint_url": "https://mint.test",
|
||||
}
|
||||
)
|
||||
return row
|
||||
|
||||
|
||||
class _Session:
|
||||
def __init__(self, transaction: CashuTransaction) -> None:
|
||||
self.transaction = transaction
|
||||
self.commit = AsyncMock()
|
||||
|
||||
async def __aenter__(self) -> "_Session":
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
async def exec(self, query: object) -> MagicMock:
|
||||
result = MagicMock()
|
||||
result.first.return_value = self.transaction
|
||||
return result
|
||||
|
||||
def add(self, transaction: CashuTransaction) -> None:
|
||||
self.transaction = transaction
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_topup_persists_before_sending_and_marks_success_collected() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=0)
|
||||
provider.topup = AsyncMock(return_value={"balance": 50})
|
||||
transaction = CashuTransaction(
|
||||
token="cashu-token", amount=50, unit="sat", source="auto_topup"
|
||||
)
|
||||
session = _Session(transaction)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.send_token",
|
||||
AsyncMock(return_value="cashu-token"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.store_cashu_transaction",
|
||||
AsyncMock(return_value=True),
|
||||
) as store,
|
||||
patch("routstr.upstream.auto_topup.create_session", return_value=session),
|
||||
):
|
||||
await _check_and_topup(_row())
|
||||
|
||||
store.assert_awaited_once_with(
|
||||
token="cashu-token",
|
||||
amount=50,
|
||||
unit="sat",
|
||||
mint_url="https://mint.test",
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
)
|
||||
provider.topup.assert_awaited_once_with("cashu-token")
|
||||
assert transaction.collected is True
|
||||
session.commit.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("outcome", [{"error": "rejected"}, RuntimeError("network")])
|
||||
async def test_auto_topup_failure_leaves_persisted_token_uncollected(
|
||||
outcome: object,
|
||||
) -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=0)
|
||||
provider.topup = AsyncMock(
|
||||
side_effect=outcome if isinstance(outcome, Exception) else None,
|
||||
return_value=outcome,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.send_token",
|
||||
AsyncMock(return_value="cashu-token"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.store_cashu_transaction",
|
||||
AsyncMock(return_value=True),
|
||||
),
|
||||
patch("routstr.upstream.auto_topup.create_session") as create_session,
|
||||
):
|
||||
if isinstance(outcome, Exception):
|
||||
with pytest.raises(RuntimeError):
|
||||
await _check_and_topup(_row())
|
||||
else:
|
||||
await _check_and_topup(_row())
|
||||
|
||||
create_session.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_topup_does_not_send_untracked_token() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=0)
|
||||
provider.topup = AsyncMock()
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.send_token",
|
||||
AsyncMock(return_value="cashu-token"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.store_cashu_transaction",
|
||||
AsyncMock(side_effect=RuntimeError("database unavailable")),
|
||||
),
|
||||
):
|
||||
await _check_and_topup(_row())
|
||||
provider.topup.assert_not_awaited()
|
||||
@@ -0,0 +1,56 @@
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.core.db import store_cashu_transaction
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"error",
|
||||
[
|
||||
OSError("disk full"),
|
||||
RuntimeError("connection lost"),
|
||||
ConnectionRefusedError("database unavailable"),
|
||||
],
|
||||
)
|
||||
async def test_store_cashu_transaction_propagates_commit_errors(
|
||||
error: Exception,
|
||||
) -> None:
|
||||
session = AsyncMock()
|
||||
session.commit.side_effect = error
|
||||
session.__aenter__.return_value = session
|
||||
session.__aexit__.return_value = None
|
||||
|
||||
with (
|
||||
patch("routstr.core.db.create_session", return_value=session),
|
||||
patch("routstr.core.db.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(type(error), match=str(error)):
|
||||
await store_cashu_transaction(
|
||||
token="cashuAtest",
|
||||
amount=1_000,
|
||||
unit="sat",
|
||||
mint_url="https://mint.example",
|
||||
typ="out",
|
||||
request_id="request-1",
|
||||
)
|
||||
|
||||
critical.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_cashu_transaction_returns_true_after_commit() -> None:
|
||||
session = AsyncMock()
|
||||
session.__aenter__.return_value = session
|
||||
session.__aexit__.return_value = None
|
||||
|
||||
with patch("routstr.core.db.create_session", return_value=session):
|
||||
stored = await store_cashu_transaction(
|
||||
token="cashuAtest",
|
||||
amount=1_000,
|
||||
unit="sat",
|
||||
)
|
||||
|
||||
assert stored is True
|
||||
session.commit.assert_awaited_once()
|
||||
@@ -0,0 +1,91 @@
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core import db
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cashu_transaction_storage_retries_then_succeeds() -> None:
|
||||
store = AsyncMock(side_effect=[OSError("database locked"), True])
|
||||
sleep = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.core.db.store_cashu_transaction", store),
|
||||
patch("routstr.core.db.asyncio.sleep", sleep),
|
||||
):
|
||||
stored = await db.store_cashu_transaction_with_retry(
|
||||
token="cashuAretry",
|
||||
amount=100,
|
||||
unit="sat",
|
||||
)
|
||||
|
||||
assert stored is True
|
||||
assert store.await_count == 2
|
||||
sleep.assert_awaited_once_with(0.25)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cashu_transaction_retry_is_idempotent_after_ambiguous_commit() -> None:
|
||||
engine = create_async_engine("sqlite+aiosqlite://")
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
original_store = db.store_cashu_transaction
|
||||
attempts = 0
|
||||
|
||||
async def ambiguous_store(**kwargs: Any) -> bool:
|
||||
nonlocal attempts
|
||||
attempts += 1
|
||||
stored = await original_store(**kwargs)
|
||||
if attempts == 1:
|
||||
raise OSError("connection dropped after commit")
|
||||
return stored
|
||||
|
||||
with (
|
||||
patch.object(db, "engine", engine),
|
||||
patch("routstr.core.db.store_cashu_transaction", ambiguous_store),
|
||||
patch("routstr.core.db.asyncio.sleep", AsyncMock()),
|
||||
):
|
||||
stored = await db.store_cashu_transaction_with_retry(
|
||||
token="cashuAambiguous",
|
||||
amount=100,
|
||||
unit="sat",
|
||||
)
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
result = await session.exec(select(db.CashuTransaction))
|
||||
transactions = result.all()
|
||||
|
||||
assert stored is True
|
||||
assert attempts == 2
|
||||
assert len(transactions) == 1
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cashu_transaction_storage_raises_after_bounded_retries() -> None:
|
||||
error = OSError("database unavailable")
|
||||
store = AsyncMock(side_effect=error)
|
||||
sleep = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.core.db.store_cashu_transaction", store),
|
||||
patch("routstr.core.db.asyncio.sleep", sleep),
|
||||
patch("routstr.core.db.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(OSError, match="database unavailable"):
|
||||
await db.store_cashu_transaction_with_retry(
|
||||
token="cashuAfail",
|
||||
amount=100,
|
||||
unit="sat",
|
||||
max_attempts=3,
|
||||
)
|
||||
|
||||
assert store.await_count == 3
|
||||
assert [call.args[0] for call in sleep.await_args_list] == [0.25, 0.5]
|
||||
critical.assert_called_once()
|
||||
@@ -465,9 +465,179 @@ async def test_cache_read_only_usd_cost_response_is_billed(
|
||||
assert result.cache_read_input_tokens == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("total_cost", "input_cost", "output_cost", "expected_msats"),
|
||||
[
|
||||
(0.000471, 0.00023451, 0.00023649, 9420),
|
||||
(0.00000004, 0.00000002, 0.00000002, 1),
|
||||
],
|
||||
)
|
||||
async def test_small_usd_cost_components_sum_to_rounded_total(
|
||||
total_cost: float,
|
||||
input_cost: float,
|
||||
output_cost: float,
|
||||
expected_msats: int,
|
||||
) -> None:
|
||||
"""Small USD component costs must retain every billed millisatoshi."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 1,
|
||||
"cost_details": {
|
||||
"total_cost": total_cost,
|
||||
"input_cost": input_cost,
|
||||
"output_cost": output_cost,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = await calculate_cost(response, max_cost=100000)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.total_msats == expected_msats
|
||||
assert result.input_msats + result.output_msats == result.total_msats
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openrouter_upstream_inference_cost_components_are_used() -> None:
|
||||
"""OpenRouter component aliases must determine the input/output split."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 375,
|
||||
"completion_tokens": 158,
|
||||
"total_tokens": 533,
|
||||
"cost": 0.00022354,
|
||||
"is_byok": False,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 286,
|
||||
"cache_write_tokens": 0,
|
||||
},
|
||||
"cost_details": {
|
||||
"upstream_inference_cost": 0.00022354,
|
||||
"upstream_inference_prompt_cost": 0.00004974,
|
||||
"upstream_inference_completions_cost": 0.0001738,
|
||||
},
|
||||
"completion_tokens_details": {"reasoning_tokens": 17},
|
||||
},
|
||||
}
|
||||
|
||||
result = await calculate_cost(response, max_cost=100000)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_msats == 994
|
||||
assert result.output_msats == 3477
|
||||
assert result.input_msats + result.output_msats == result.total_msats == 4471
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# PPQ.AI BYOK: upstream_inference_cost + BYOK fee billing
|
||||
#
|
||||
# PPQ.AI (bring-your-own-key) returns a small ~5 % BYOK routing fee in
|
||||
# ``usage.cost`` and the real inference cost in
|
||||
# ``cost_details.upstream_inference_cost``. The old code billed only the fee,
|
||||
# under-charging by ~20×. The fix bills ``upstream_inference_cost + byok_fee``
|
||||
# — what PPQ actually deducts from the balance.
|
||||
# Payload numbers are from a live ``glm-5.2-fast`` request (GitHub issue #615).
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_byok_bills_upstream_inference_cost_plus_fee() -> None:
|
||||
"""PPQ.AI BYOK must bill upstream_inference_cost + byok_fee, not just the
|
||||
fee. Mirrors the live request from GitHub issue #615."""
|
||||
response = {
|
||||
"model": "glm-5.2-fast",
|
||||
"usage": {
|
||||
"prompt_tokens": 164371, # includes 159301 cached
|
||||
"completion_tokens": 99,
|
||||
"cost": 0.002260057305, # ~5% BYOK routing fee
|
||||
"is_byok": True,
|
||||
"prompt_tokens_details": {"cached_tokens": 159301},
|
||||
"cost_details": {
|
||||
"upstream_inference_cost": 0.04475361,
|
||||
"upstream_inference_prompt_cost": 0.04410021,
|
||||
"upstream_inference_completions_cost": 0.0006534,
|
||||
},
|
||||
},
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
# The fix bills upstream_inference_cost + byok_fee (~0.047 USD → ~940k
|
||||
# msats), not the fee alone (~0.0023 USD → ~45k msats). ~20× correction.
|
||||
assert result.total_msats == 940274
|
||||
assert result.input_msats + result.output_msats == result.total_msats
|
||||
assert result.input_msats == 926546
|
||||
assert result.output_msats == 13728
|
||||
assert result.total_usd == pytest.approx(0.047013667305)
|
||||
# Token normalisation (OpenAI dialect: cached included in prompt_tokens)
|
||||
assert result.input_tokens == 5070 # 164371 - 159301
|
||||
assert result.cache_read_input_tokens == 159301
|
||||
assert result.output_tokens == 99
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_byok_fee_only_would_undercharge() -> None:
|
||||
"""Sanity check: billing only usage.cost (the BYOK fee) under-charges by
|
||||
~20×. This documents the regression the fix prevents."""
|
||||
response = {
|
||||
"model": "glm-5.2-fast",
|
||||
"usage": {
|
||||
"prompt_tokens": 164371,
|
||||
"completion_tokens": 99,
|
||||
"cost": 0.002260057305, # BYOK fee only — no upstream_inference_cost
|
||||
"is_byok": True,
|
||||
"prompt_tokens_details": {"cached_tokens": 159301},
|
||||
},
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
# Without upstream_inference_cost, only the fee is billed — the old bug.
|
||||
assert result.total_msats == 45202
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 13: Missing Usage Block
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_upstream_cost_uses_litellm_model_pricing(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Token usage without an upstream cost is priced from LiteLLM."""
|
||||
monkeypatch.setattr(settings, "fixed_pricing", True)
|
||||
monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0)
|
||||
monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0)
|
||||
monkeypatch.setattr(
|
||||
"routstr.payment.models.litellm_cost_entry",
|
||||
lambda model: {
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000002,
|
||||
},
|
||||
)
|
||||
response = {
|
||||
"model": "priced-by-litellm",
|
||||
"usage": {
|
||||
"prompt_tokens": 90,
|
||||
"completion_tokens": 80,
|
||||
"total_tokens": 170,
|
||||
"prompt_tokens_details": {"cached_tokens": 0},
|
||||
"completion_tokens_details": {"reasoning_tokens": 74},
|
||||
"prompt_cache_hit_tokens": 0,
|
||||
"prompt_cache_miss_tokens": 90,
|
||||
},
|
||||
}
|
||||
|
||||
result = await calculate_cost(response, max_cost=10000)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert not isinstance(result, MaxCostData)
|
||||
assert result.input_msats == 1800
|
||||
assert result.output_msats == 3200
|
||||
assert result.input_msats + result.output_msats == result.total_msats == 5000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_usage_block(mock_fixed_pricing: None) -> None:
|
||||
"""When usage is missing, return MaxCostData with zero tokens."""
|
||||
|
||||
@@ -0,0 +1,173 @@
|
||||
"""Coverage tests for admin.py (currently 35%).
|
||||
|
||||
Tests admin endpoints that are testable without full app setup:
|
||||
withdraw validation, password update, CLI token lifecycle.
|
||||
"""
|
||||
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
# ===========================================================================
|
||||
# withdraw — validation and edge cases
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_withdraw_rejects_zero_amount() -> None:
|
||||
"""withdraw validation rejects amount <= 0."""
|
||||
from routstr.core.admin import WithdrawRequest, withdraw
|
||||
|
||||
request = Request(scope={"type": "http", "method": "POST"})
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await withdraw(request, WithdrawRequest(amount=0, unit="sat"))
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_withdraw_rejects_negative_amount() -> None:
|
||||
"""withdraw validation rejects negative amounts."""
|
||||
from routstr.core.admin import WithdrawRequest, withdraw
|
||||
|
||||
request = Request(scope={"type": "http", "method": "POST"})
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await withdraw(request, WithdrawRequest(amount=-100, unit="sat"))
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_withdraw_rejects_insufficient_balance() -> None:
|
||||
"""withdraw returns 400 when wallet balance is insufficient."""
|
||||
from routstr.core.admin import WithdrawRequest, withdraw
|
||||
|
||||
request = Request(scope={"type": "http", "method": "POST"})
|
||||
|
||||
with patch("routstr.core.admin.get_wallet") as mock_wallet, \
|
||||
patch("routstr.core.admin.get_proofs_per_mint_and_unit") as mock_proofs, \
|
||||
patch("routstr.core.admin.slow_filter_spend_proofs") as mock_filter:
|
||||
|
||||
mock_w = Mock()
|
||||
mock_w.keysets = {}
|
||||
mock_w.proofs = []
|
||||
mock_wallet.return_value = mock_w
|
||||
mock_proofs.return_value = []
|
||||
mock_filter.return_value = []
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await withdraw(request, WithdrawRequest(amount=1000000, unit="sat"))
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Insufficient" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# update_password — validation
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_password_rejects_empty_new() -> None:
|
||||
"""update_password rejects empty new password."""
|
||||
from routstr.core.admin import PasswordUpdate, update_password
|
||||
|
||||
request = Request(scope={"type": "http", "method": "POST"})
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await update_password(
|
||||
request,
|
||||
PasswordUpdate(current_password="old", new_password=""),
|
||||
)
|
||||
|
||||
# Returns 500 (no admin password configured) or 400 (validation)
|
||||
assert exc_info.value.status_code in (400, 500, 422)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_password_rejects_short_new() -> None:
|
||||
"""update_password rejects short passwords."""
|
||||
from routstr.core.admin import PasswordUpdate, update_password
|
||||
|
||||
request = Request(scope={"type": "http", "method": "POST"})
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await update_password(
|
||||
request,
|
||||
PasswordUpdate(current_password="old", new_password="ab"),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code in (400, 500, 422)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# require_admin_api guard
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_require_admin_rejects_no_session() -> None:
|
||||
"""require_admin_api rejects requests without admin session cookie."""
|
||||
from routstr.core.admin import require_admin_api
|
||||
|
||||
request = Request(scope={
|
||||
"type": "http",
|
||||
"method": "GET",
|
||||
"headers": [],
|
||||
})
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await require_admin_api(request)
|
||||
|
||||
# 401 or 403 depending on auth configuration
|
||||
assert exc_info.value.status_code in (401, 403)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# _validate_slug
|
||||
# ===========================================================================
|
||||
|
||||
def test_validate_slug_accepts_valid() -> None:
|
||||
"""Valid slugs pass validation."""
|
||||
from routstr.core.admin import _validate_slug
|
||||
|
||||
assert _validate_slug("valid-slug") == "valid-slug"
|
||||
assert _validate_slug("valid123") == "valid123"
|
||||
assert _validate_slug("my-provider") == "my-provider"
|
||||
|
||||
|
||||
def test_validate_slug_rejects_spaces() -> None:
|
||||
"""Slugs with spaces are rejected."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.core.admin import _validate_slug
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
_validate_slug("invalid slug")
|
||||
|
||||
|
||||
def test_validate_slug_rejects_too_short() -> None:
|
||||
"""Slugs shorter than 3 chars are rejected."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.core.admin import _validate_slug
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
_validate_slug("ab")
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# admin login endpoint
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_login_requires_payload() -> None:
|
||||
"""admin_login requires a payload — verify it exists."""
|
||||
# Verify the function signature
|
||||
import inspect
|
||||
|
||||
from routstr.core.admin import admin_login
|
||||
sig = inspect.signature(admin_login)
|
||||
params = list(sig.parameters.keys())
|
||||
assert "request" in params
|
||||
assert "payload" in params or len(params) >= 2
|
||||
@@ -0,0 +1,204 @@
|
||||
"""Coverage tests for base.py (currently 41%).
|
||||
|
||||
Tests preparers, builders, accessors, and model cache methods.
|
||||
"""
|
||||
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
# ===========================================================================
|
||||
# prepare_headers
|
||||
# ===========================================================================
|
||||
|
||||
def test_prepare_headers_adds_auth() -> None:
|
||||
"""API key is added as Bearer token."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
headers = p.prepare_headers({})
|
||||
|
||||
assert "Authorization" in headers
|
||||
assert headers["Authorization"] == "Bearer sk-test-key"
|
||||
|
||||
|
||||
def test_prepare_headers_preserves_existing() -> None:
|
||||
"""Existing headers are preserved."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
headers = p.prepare_headers({"X-Custom": "value", "Content-Type": "application/json"})
|
||||
|
||||
assert headers["X-Custom"] == "value"
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
|
||||
|
||||
def test_prepare_headers_auth_header_passthrough() -> None:
|
||||
"""Authorization header is handled — verify current behaviour."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
headers = p.prepare_headers({"Authorization": "Bearer user-key"})
|
||||
|
||||
# Currently provider key is used (may be intentional for proxy pattern)
|
||||
assert "Authorization" in headers
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# prepare_params
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prepare_params_passes_through() -> None:
|
||||
"""Query params are preserved by default."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
params = p.prepare_params("/v1/chat/completions", {"temperature": "0.7"})
|
||||
|
||||
assert params["temperature"] == "0.7"
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# transform_model_name / normalize_request_path / get_request_base_url
|
||||
# ===========================================================================
|
||||
|
||||
def test_transform_model_name_default_passthrough() -> None:
|
||||
"""Default returns model_id unchanged."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
assert p.transform_model_name("gpt-4") == "gpt-4"
|
||||
assert p.transform_model_name("") == ""
|
||||
|
||||
|
||||
def test_normalize_request_path_passthrough() -> None:
|
||||
"""Default returns path unchanged."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
assert p.normalize_request_path("/v1/chat/completions") == "/v1/chat/completions"
|
||||
|
||||
|
||||
def test_get_request_base_url_default() -> None:
|
||||
"""Default returns the provider's base_url."""
|
||||
p = BaseUpstreamProvider("https://api.test.com/v1", "sk-test-key")
|
||||
url = p.get_request_base_url("/v1/chat/completions")
|
||||
assert url == "https://api.test.com/v1"
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# build_request_url
|
||||
# ===========================================================================
|
||||
|
||||
def test_build_request_url_combines_base_and_path() -> None:
|
||||
"""Combines base_url and path."""
|
||||
p = BaseUpstreamProvider("https://api.test.com/v1", "sk-test-key")
|
||||
url = p.build_request_url("/chat/completions")
|
||||
assert "api.test.com" in url
|
||||
assert "/chat/completions" in url
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# get_litellm_provider_prefix / get_provider_metadata
|
||||
# ===========================================================================
|
||||
|
||||
def test_get_litellm_provider_prefix_default() -> None:
|
||||
"""Default returns a string prefix."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
prefix = p.get_litellm_provider_prefix()
|
||||
assert isinstance(prefix, str)
|
||||
|
||||
|
||||
def test_get_provider_metadata_returns_dict() -> None:
|
||||
"""Default metadata has name and capabilities."""
|
||||
metadata = BaseUpstreamProvider.get_provider_metadata()
|
||||
assert isinstance(metadata, dict)
|
||||
assert "name" in metadata
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# from_db_row
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_from_db_row_returns_provider() -> None:
|
||||
"""from_db_row constructs a provider from a valid row."""
|
||||
mock_row = Mock()
|
||||
mock_row.base_url = "https://api.test.com"
|
||||
mock_row.api_key = "sk-test-key"
|
||||
mock_row.slug = "test-slug"
|
||||
mock_row.provider_fee = 1.0
|
||||
mock_row.field_overrides = None
|
||||
mock_row.name = "Test"
|
||||
|
||||
result = BaseUpstreamProvider.from_db_row(mock_row)
|
||||
assert result is not None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# prepare_request_body
|
||||
# ===========================================================================
|
||||
|
||||
def test_prepare_request_body_with_model() -> None:
|
||||
"""prepare_request_body takes bytes body and Model object."""
|
||||
mock_model = Mock()
|
||||
mock_model.id = "gpt-4"
|
||||
mock_model.forwarded_model_id = None
|
||||
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
|
||||
# None body returns None
|
||||
result = p.prepare_request_body(None, mock_model)
|
||||
assert result is None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# prepare_responses_request_body
|
||||
# ===========================================================================
|
||||
|
||||
def test_prepare_responses_request_body_none() -> None:
|
||||
"""None body returns None."""
|
||||
model_obj = Mock()
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
result = p.prepare_responses_request_body(None, model_obj)
|
||||
assert result is None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# _upstream_accepts_cache_control
|
||||
# ===========================================================================
|
||||
|
||||
def test_upstream_accepts_cache_control_default() -> None:
|
||||
"""Default: upstream does NOT accept cache-control."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
assert p._upstream_accepts_cache_control() is False
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# inject_cost_metadata
|
||||
# ===========================================================================
|
||||
|
||||
def test_inject_cost_metadata_adds_metadata() -> None:
|
||||
"""Cost metadata is injected into the response dict."""
|
||||
mock_key = Mock()
|
||||
mock_key.balance_msat = 500000
|
||||
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
data = {"model": "gpt-4", "usage": {"prompt_tokens": 100}}
|
||||
cost_data = {
|
||||
"base_msats": 200000,
|
||||
"input_msats": 100000,
|
||||
"output_msats": 100000,
|
||||
"total_msats": 200000,
|
||||
"total_usd": 0.01,
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
}
|
||||
|
||||
p.inject_cost_metadata(data, cost_data, mock_key)
|
||||
|
||||
# Metadata is nested under metadata.routstr.cost
|
||||
assert "metadata" in data or "routstr_cost" in data or "cost" in data
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# _apply_provider_field
|
||||
# ===========================================================================
|
||||
|
||||
def test_apply_provider_field_adds_to_response() -> None:
|
||||
"""Provider field is added to response JSON."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test-key")
|
||||
data = {"id": "chatcmpl-123"}
|
||||
p._apply_provider_field(data)
|
||||
assert "provider" in data
|
||||
@@ -0,0 +1,219 @@
|
||||
"""Additional coverage tests for base.py (41% → target 50%+).
|
||||
|
||||
Tests error message extraction, static helpers, model cache, and cost hooks.
|
||||
|
||||
These test existing correct behavior — all should PASS.
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
# ===========================================================================
|
||||
# _extract_upstream_error_message
|
||||
# ===========================================================================
|
||||
|
||||
def test_extract_error_from_json_body() -> None:
|
||||
"""Error message is extracted from JSON upstream error response."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
body = json.dumps({"error": {"message": "Model not found", "type": "not_found"}}).encode()
|
||||
|
||||
msg, error_type = p._extract_upstream_error_message(body)
|
||||
|
||||
assert "Model not found" in msg
|
||||
assert error_type == "not_found"
|
||||
|
||||
|
||||
def test_extract_error_from_simple_json() -> None:
|
||||
"""Simple JSON error with direct message key."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
body = json.dumps({"message": "Rate limit exceeded"}).encode()
|
||||
|
||||
msg, error_type = p._extract_upstream_error_message(body)
|
||||
|
||||
assert "Rate limit" in msg
|
||||
|
||||
|
||||
def test_extract_error_from_text_body() -> None:
|
||||
"""Non-JSON text body is returned as-is."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
|
||||
msg, error_type = p._extract_upstream_error_message(b"Internal Server Error")
|
||||
|
||||
assert "Internal Server Error" in msg
|
||||
|
||||
|
||||
def test_extract_error_empty_body() -> None:
|
||||
"""Empty body returns a generic message."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
|
||||
msg, error_type = p._extract_upstream_error_message(b"")
|
||||
|
||||
assert isinstance(msg, str)
|
||||
assert len(msg) > 0
|
||||
|
||||
|
||||
def test_extract_error_simple_error_string_not_parsed() -> None:
|
||||
"""JSON error as plain string (not dict) falls through to generic message."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
body = json.dumps({"error": "Invalid API key"}).encode()
|
||||
|
||||
msg, error_type = p._extract_upstream_error_message(body)
|
||||
|
||||
# Simple error strings not nested in a dict object use generic message
|
||||
assert "Upstream request failed" in msg or "Invalid" in msg
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# on_upstream_error_redirect
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_upstream_error_redirect_noop() -> None:
|
||||
"""Default implementation is a no-op for non-redirect statuses."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
result = await p.on_upstream_error_redirect(402, "Insufficient balance")
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_on_upstream_error_redirect_429() -> None:
|
||||
"""429 rate limit passes through (subclasses may override)."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
result = await p.on_upstream_error_redirect(429, "Rate limited")
|
||||
assert result is None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# _fold_cache_into_input_tokens (static method)
|
||||
# ===========================================================================
|
||||
|
||||
def test_fold_cache_no_cache_data() -> None:
|
||||
"""Usage without cache details is unchanged."""
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
usage = Mock()
|
||||
usage.prompt_tokens = 100
|
||||
del usage.prompt_tokens_details # No cache details
|
||||
|
||||
BaseUpstreamProvider._fold_cache_into_input_tokens(usage)
|
||||
# Should not modify the usage object when no cache exists
|
||||
|
||||
|
||||
def test_fold_cache_preserves_total() -> None:
|
||||
"""Total prompt tokens remain the same after folding cache."""
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
usage = Mock()
|
||||
usage.prompt_tokens = 100
|
||||
details = Mock()
|
||||
details.cached_tokens = 30
|
||||
usage.prompt_tokens_details = details
|
||||
|
||||
BaseUpstreamProvider._fold_cache_into_input_tokens(usage)
|
||||
# prompt_tokens should still be 100 (total unchanged)
|
||||
assert usage.prompt_tokens == 100
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# get_cached_models / get_cached_model_by_id
|
||||
# ===========================================================================
|
||||
|
||||
def test_get_cached_models_returns_list() -> None:
|
||||
"""get_cached_models always returns a list."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
models = p.get_cached_models()
|
||||
assert isinstance(models, list)
|
||||
|
||||
|
||||
def test_get_cached_model_by_id_unknown_returns_none() -> None:
|
||||
"""Unknown model ID returns None."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
result = p.get_cached_model_by_id("nonexistent-model-xyz-12345")
|
||||
assert result is None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# get_x_cashu_cost
|
||||
# ===========================================================================
|
||||
|
||||
def test_get_x_cashu_cost_with_usage() -> None:
|
||||
"""Cost is calculated from response data with usage info."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
response_data = {
|
||||
"model": "gpt-4",
|
||||
"usage": {"prompt_tokens": 100, "completion_tokens": 50},
|
||||
}
|
||||
|
||||
result = p.get_x_cashu_cost(response_data, 100000, None)
|
||||
|
||||
# Either returns None (needs more data) or a cost object
|
||||
assert result is not None
|
||||
|
||||
|
||||
def test_get_x_cashu_cost_no_usage() -> None:
|
||||
"""Response without usage returns MaxCostData."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
response_data = {"model": "gpt-4"}
|
||||
|
||||
result = p.get_x_cashu_cost(response_data, 100000, None)
|
||||
|
||||
# Without usage, uses max_cost
|
||||
assert result is not None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# get_balance
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_balance_raises_not_implemented() -> None:
|
||||
"""Default get_balance raises NotImplementedError (no account support)."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
with pytest.raises(NotImplementedError):
|
||||
await p.get_balance()
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# refresh_models_cache
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_models_cache_no_providers() -> None:
|
||||
"""refresh_models_cache handles empty provider list gracefully."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
# Default implementation may be a no-op or raise
|
||||
try:
|
||||
await p.refresh_models_cache()
|
||||
except Exception:
|
||||
pass # May fail without DB — that's fine
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# fetch_models
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_models_returns_list() -> None:
|
||||
"""fetch_models returns a model list (or empty) for default provider."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
try:
|
||||
result = await p.fetch_models()
|
||||
assert isinstance(result, list)
|
||||
except Exception:
|
||||
pass # May fail without network
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# create_account
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_account_raises_not_implemented() -> None:
|
||||
"""Default create_account raises NotImplementedError."""
|
||||
p = BaseUpstreamProvider("https://api.test.com", "sk-test")
|
||||
with pytest.raises(NotImplementedError):
|
||||
await p.create_account()
|
||||
@@ -0,0 +1,150 @@
|
||||
"""Coverage-filling tests for middleware.py (currently 38% coverage).
|
||||
|
||||
Only LoggingMiddleware and request_id_context exist on main.
|
||||
ConcurrencyLimiterMiddleware + TimeoutMiddleware are on an unmerged branch.
|
||||
"""
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LoggingMiddleware
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_logging_middleware_adds_request_id() -> None:
|
||||
"""Every request gets an x-routstr-request-id header."""
|
||||
from routstr.core.middleware import LoggingMiddleware
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
@app.get("/test")
|
||||
async def test_endpoint(request: Request) -> dict:
|
||||
assert hasattr(request.state, "request_id")
|
||||
assert request.state.request_id is not None
|
||||
return {"ok": True}
|
||||
|
||||
app.add_middleware(LoggingMiddleware)
|
||||
|
||||
client = TestClient(app)
|
||||
response = client.get("/test")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert "x-routstr-request-id" in response.headers
|
||||
assert len(response.headers["x-routstr-request-id"]) == 36 # UUID4 length
|
||||
|
||||
|
||||
def test_logging_middleware_skips_head_requests() -> None:
|
||||
"""HEAD requests are skipped by _should_log (health probes)."""
|
||||
from routstr.core.middleware import LoggingMiddleware
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
@app.head("/test")
|
||||
async def test_endpoint(request: Request) -> dict:
|
||||
return {"ok": True}
|
||||
|
||||
app.add_middleware(LoggingMiddleware)
|
||||
|
||||
client = TestClient(app)
|
||||
response = client.head("/test")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert "x-routstr-request-id" in response.headers
|
||||
|
||||
|
||||
def test_logging_middleware_skips_options_requests() -> None:
|
||||
"""OPTIONS requests (CORS preflight) are skipped."""
|
||||
from routstr.core.middleware import LoggingMiddleware
|
||||
|
||||
app = FastAPI()
|
||||
|
||||
@app.options("/test")
|
||||
async def test_endpoint(request: Request) -> dict:
|
||||
return {"ok": True}
|
||||
|
||||
app.add_middleware(LoggingMiddleware)
|
||||
|
||||
client = TestClient(app)
|
||||
response = client.options("/test")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert "x-routstr-request-id" in response.headers
|
||||
|
||||
|
||||
def test_should_log_rejects_admin_api_prefix() -> None:
|
||||
"""Admin API polling paths are skipped."""
|
||||
from routstr.core.middleware import _should_log
|
||||
|
||||
assert _should_log("GET", "/admin/api/balances") is False
|
||||
assert _should_log("GET", "/admin/api/logs") is False
|
||||
assert _should_log("GET", "/admin/api/providers") is False
|
||||
|
||||
|
||||
def test_should_log_rejects_nextjs_chunks() -> None:
|
||||
"""Next.js static chunks are skipped."""
|
||||
from routstr.core.middleware import _should_log
|
||||
|
||||
assert _should_log("GET", "/_next/static/chunks/main.js") is False
|
||||
assert _should_log("GET", "/_next/data/build-id/page.json") is False
|
||||
|
||||
|
||||
def test_should_log_rejects_exact_paths() -> None:
|
||||
"""Exact paths like /favicon.ico are skipped."""
|
||||
from routstr.core.middleware import _should_log
|
||||
|
||||
assert _should_log("GET", "/favicon.ico") is False
|
||||
assert _should_log("GET", "/v1/wallet/info") is False
|
||||
assert _should_log("GET", "/index.txt") is False
|
||||
assert _should_log("GET", "/login/index.txt") is False
|
||||
|
||||
|
||||
def test_should_log_accepts_normal_paths() -> None:
|
||||
"""Normal API paths are logged."""
|
||||
from routstr.core.middleware import _should_log
|
||||
|
||||
assert _should_log("GET", "/v1/chat/completions") is True
|
||||
assert _should_log("POST", "/v1/chat/completions") is True
|
||||
assert _should_log("GET", "/v1/models") is True
|
||||
assert _should_log("POST", "/api/some-endpoint") is True
|
||||
|
||||
|
||||
def test_should_log_accepts_non_skipped_path() -> None:
|
||||
"""Generic paths not in skip list are logged."""
|
||||
from routstr.core.middleware import _should_log
|
||||
|
||||
assert _should_log("GET", "/some/random/path") is True
|
||||
assert _should_log("POST", "/api/custom") is True
|
||||
|
||||
|
||||
def test_request_id_context_is_contextvar() -> None:
|
||||
"""request_id_context is a ContextVar[str | None] with no default value."""
|
||||
from contextvars import ContextVar
|
||||
|
||||
from routstr.core.middleware import request_id_context
|
||||
|
||||
assert isinstance(request_id_context, ContextVar)
|
||||
# ContextVar without a default raises LookupError when accessed without being set
|
||||
try:
|
||||
val = request_id_context.get()
|
||||
# If it returns, it should be None
|
||||
assert val is None
|
||||
except LookupError:
|
||||
# Expected: ContextVar with no default raises LookupError
|
||||
pass
|
||||
|
||||
|
||||
def test_middleware_exports() -> None:
|
||||
"""Only LoggingMiddleware is exported on main."""
|
||||
from routstr.core.middleware import LoggingMiddleware, request_id_context
|
||||
|
||||
assert LoggingMiddleware is not None
|
||||
assert request_id_context is not None
|
||||
|
||||
|
||||
def test_middleware_skips_health_probe_path() -> None:
|
||||
"""Health probe paths pass through without logging."""
|
||||
from routstr.core.middleware import _should_log
|
||||
|
||||
# HEAD method is always skipped regardless of path
|
||||
assert _should_log("HEAD", "/v1/chat/completions") is False
|
||||
assert _should_log("OPTIONS", "/v1/chat/completions") is False
|
||||
@@ -0,0 +1,181 @@
|
||||
"""Coverage-filling tests for payment/helpers.py (currently 52% coverage).
|
||||
|
||||
Tests the real public API: check_token_balance, get_max_cost_for_model,
|
||||
estimate_tokens, create_error_response, etc.
|
||||
"""
|
||||
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# check_token_balance
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_token_balance_x_cashu_present() -> None:
|
||||
"""X-Cashu header triggers token deserialization and balance check."""
|
||||
from routstr.payment.helpers import check_token_balance
|
||||
|
||||
headers = {"x-cashu": "cashuAtest_token"}
|
||||
body = {"model": "gpt-4"}
|
||||
|
||||
with patch("routstr.payment.helpers.deserialize_token_from_string") as mock_deser:
|
||||
mock_token = Mock()
|
||||
mock_token.amount = 50000
|
||||
mock_token.unit = "sat"
|
||||
mock_deser.return_value = mock_token
|
||||
|
||||
# Should not raise — balance is sufficient
|
||||
check_token_balance(headers, body, 1000)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_token_balance_no_x_cashu_raises() -> None:
|
||||
"""Missing X-Cashu header raises HTTPException (401 on main)."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.payment.helpers import check_token_balance
|
||||
|
||||
headers = {}
|
||||
body = {"model": "gpt-4"}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
check_token_balance(headers, body, 1000)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_token_balance_insufficient_raises() -> None:
|
||||
"""Token with insufficient balance raises HTTPException 402.
|
||||
|
||||
max_cost_for_model is in msat, so with amount=100 sat (=100,000 msat),
|
||||
max_cost=200,000 msat triggers the insufficient balance check.
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.payment.helpers import check_token_balance
|
||||
|
||||
headers = {"x-cashu": "cashuAtest_token"}
|
||||
body = {"model": "gpt-4"}
|
||||
|
||||
with patch("routstr.payment.helpers.deserialize_token_from_string") as mock_deser:
|
||||
mock_token = Mock()
|
||||
mock_token.amount = 100 # 100 sat
|
||||
mock_token.unit = "sat"
|
||||
mock_deser.return_value = mock_token
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
# 200,000 msat > 100,000 msat (100 sat * 1000)
|
||||
check_token_balance(headers, body, 200000)
|
||||
|
||||
assert exc_info.value.status_code == 402
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# estimate_tokens
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_estimate_tokens_empty_messages() -> None:
|
||||
"""Empty message list returns 0 tokens."""
|
||||
from routstr.payment.helpers import estimate_tokens
|
||||
|
||||
result = estimate_tokens([])
|
||||
|
||||
assert result == 0
|
||||
|
||||
|
||||
def test_estimate_tokens_text_content() -> None:
|
||||
"""Text messages are counted."""
|
||||
from routstr.payment.helpers import estimate_tokens
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
]
|
||||
|
||||
result = estimate_tokens(messages)
|
||||
|
||||
assert result > 0
|
||||
assert isinstance(result, int)
|
||||
|
||||
|
||||
def test_estimate_tokens_long_text() -> None:
|
||||
"""Longer messages produce higher token counts."""
|
||||
from routstr.payment.helpers import estimate_tokens
|
||||
|
||||
short = estimate_tokens([{"role": "user", "content": "Hi"}])
|
||||
long = estimate_tokens([{"role": "user", "content": "Hello " * 100}])
|
||||
|
||||
assert long > short
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# create_error_response
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_create_error_response_402() -> None:
|
||||
"""402 Payment Required error is properly formatted."""
|
||||
from fastapi import Request
|
||||
|
||||
from routstr.payment.helpers import create_error_response
|
||||
|
||||
request = Request(scope={"type": "http", "method": "GET"})
|
||||
result = create_error_response("insufficient_funds", "Insufficient balance", 402, request)
|
||||
|
||||
assert result.status_code == 402
|
||||
|
||||
|
||||
def test_create_error_response_500() -> None:
|
||||
"""500 Internal Server Error is properly formatted."""
|
||||
from fastapi import Request
|
||||
|
||||
from routstr.payment.helpers import create_error_response
|
||||
|
||||
request = Request(scope={"type": "http", "method": "GET"})
|
||||
result = create_error_response("server_error", "Internal error", 500, request)
|
||||
|
||||
assert result.status_code == 500
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Image token estimation helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_image_dimensions_valid_png() -> None:
|
||||
"""_get_image_dimensions returns width and height for a valid PNG."""
|
||||
from routstr.payment.helpers import _get_image_dimensions
|
||||
|
||||
# A minimal 1x1 red PNG (valid minimal file)
|
||||
png = (
|
||||
b"\x89PNG\r\n\x1a\n"
|
||||
b"\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01\x08\x02"
|
||||
b"\x00\x00\x00\x90wS\xde"
|
||||
b"\x00\x00\x00\x0cIDAT\x08\xd7c\xf8\x0f\x00\x00\x01\x01\x00\x05"
|
||||
b"\x18\xd8N"
|
||||
b"\x00\x00\x00\x00IEND\xaeB`\x82"
|
||||
)
|
||||
|
||||
w, h = _get_image_dimensions(png)
|
||||
assert w == 1
|
||||
assert h == 1
|
||||
|
||||
|
||||
def test_calculate_image_tokens_low_detail() -> None:
|
||||
"""Low detail images are always 85 tokens."""
|
||||
from routstr.payment.helpers import _calculate_image_tokens
|
||||
|
||||
tokens = _calculate_image_tokens(1024, 1024, "low")
|
||||
|
||||
assert tokens == 85
|
||||
|
||||
|
||||
def test_calculate_image_tokens_high_detail() -> None:
|
||||
"""High detail images are scaled and tile-based."""
|
||||
from routstr.payment.helpers import _calculate_image_tokens
|
||||
|
||||
tokens = _calculate_image_tokens(1024, 1024, "high")
|
||||
|
||||
assert tokens > 85
|
||||
assert isinstance(tokens, int)
|
||||
@@ -0,0 +1,161 @@
|
||||
"""Coverage tests for proxy.py (currently 47%).
|
||||
|
||||
Tests request parsing, model extraction, and routing helpers.
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
# ===========================================================================
|
||||
# parse_request_body_json
|
||||
# ===========================================================================
|
||||
|
||||
def test_parse_json_valid_body() -> None:
|
||||
"""Valid JSON body is parsed correctly for chat completions."""
|
||||
from routstr.proxy import parse_request_body_json
|
||||
|
||||
body = json.dumps({"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]}).encode()
|
||||
result = parse_request_body_json(body, "/v1/chat/completions")
|
||||
|
||||
assert result["model"] == "gpt-4"
|
||||
assert result["messages"][0]["role"] == "user"
|
||||
|
||||
|
||||
def test_parse_json_invalid_raises_400() -> None:
|
||||
"""Invalid JSON raises HTTPException 400."""
|
||||
from routstr.proxy import parse_request_body_json
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
parse_request_body_json(b"not json", "/v1/chat/completions")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
def test_parse_json_empty_body() -> None:
|
||||
"""Empty body returns empty dict."""
|
||||
from routstr.proxy import parse_request_body_json
|
||||
|
||||
result = parse_request_body_json(b"", "/v1/chat/completions")
|
||||
assert isinstance(result, dict)
|
||||
assert result == {}
|
||||
|
||||
|
||||
def test_parse_json_responses_path() -> None:
|
||||
"""Responses API path is handled."""
|
||||
from routstr.proxy import parse_request_body_json
|
||||
|
||||
body = json.dumps({"model": "gpt-4", "input": "hello"}).encode()
|
||||
result = parse_request_body_json(body, "/v1/responses")
|
||||
|
||||
assert "model" in result
|
||||
|
||||
|
||||
def test_parse_json_rejects_non_integer_max_tokens() -> None:
|
||||
"""max_tokens must be an integer."""
|
||||
from routstr.proxy import parse_request_body_json
|
||||
|
||||
body = json.dumps({"model": "gpt-4", "max_tokens": "abc"}).encode()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
parse_request_body_json(body, "/v1/chat/completions")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# extract_model_from_responses_request
|
||||
# ===========================================================================
|
||||
|
||||
def test_extract_model_from_responses() -> None:
|
||||
"""Model name is extracted from Responses API request."""
|
||||
from routstr.proxy import extract_model_from_responses_request
|
||||
|
||||
body = {"model": "gpt-4o", "input": "test"}
|
||||
model = extract_model_from_responses_request(body)
|
||||
assert model == "gpt-4o"
|
||||
|
||||
|
||||
def test_extract_model_returns_unknown_for_missing() -> None:
|
||||
"""Missing model field returns 'unknown'."""
|
||||
from routstr.proxy import extract_model_from_responses_request
|
||||
|
||||
body = {"input": "test"}
|
||||
model = extract_model_from_responses_request(body)
|
||||
assert model == "unknown"
|
||||
|
||||
|
||||
def test_extract_model_empty_body_returns_unknown() -> None:
|
||||
"""Empty body returns 'unknown'."""
|
||||
from routstr.proxy import extract_model_from_responses_request
|
||||
|
||||
model = extract_model_from_responses_request({})
|
||||
assert model == "unknown"
|
||||
|
||||
|
||||
def test_extract_model_from_input_nested() -> None:
|
||||
"""Model nested in input dict is found."""
|
||||
from routstr.proxy import extract_model_from_responses_request
|
||||
|
||||
body = {"input": {"model": "claude-sonnet", "text": "hi"}}
|
||||
model = extract_model_from_responses_request(body)
|
||||
# The function checks input_data.get("model") for nested
|
||||
assert model in ("claude-sonnet", "unknown")
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# get_model_instance / get_provider_for_model / get_unique_models
|
||||
# ===========================================================================
|
||||
|
||||
def test_get_model_instance_unknown_returns_none() -> None:
|
||||
"""Unknown model ID returns None."""
|
||||
from routstr.proxy import get_model_instance
|
||||
|
||||
result = get_model_instance("nonexistent-model-xyz-12345")
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_get_provider_for_model_unknown_returns_none() -> None:
|
||||
"""Unknown model returns None."""
|
||||
from routstr.proxy import get_provider_for_model
|
||||
|
||||
result = get_provider_for_model("nonexistent-model-xyz-12345")
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_get_unique_models_returns_list() -> None:
|
||||
"""get_unique_models always returns a list."""
|
||||
from routstr.proxy import get_unique_models
|
||||
|
||||
result = get_unique_models()
|
||||
assert isinstance(result, list)
|
||||
|
||||
|
||||
def test_get_upstreams_returns_list() -> None:
|
||||
"""get_upstreams returns a list of providers."""
|
||||
from routstr.proxy import get_upstreams
|
||||
|
||||
result = get_upstreams()
|
||||
assert isinstance(result, list)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# parse_request_body_json — nested objects
|
||||
# ===========================================================================
|
||||
|
||||
def test_parse_body_preserves_nested_objects() -> None:
|
||||
"""Nested JSON objects are preserved during parsing."""
|
||||
from routstr.proxy import parse_request_body_json
|
||||
|
||||
body = json.dumps({
|
||||
"model": "claude-3",
|
||||
"messages": [{"role": "system", "content": "You are helpful."}],
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 1024,
|
||||
}).encode()
|
||||
|
||||
result = parse_request_body_json(body, "/v1/chat/completions")
|
||||
assert result["temperature"] == 0.7
|
||||
assert result["max_tokens"] == 1024
|
||||
assert len(result["messages"]) == 1
|
||||
@@ -0,0 +1,89 @@
|
||||
"""Tests asserting CORRECT behavior for DB persistence and payout safety.
|
||||
|
||||
RED tests — FAIL against current main until bugs are fixed.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
|
||||
# ===========================================================================
|
||||
# RED TESTS: Fee payout crash safety
|
||||
# ===========================================================================
|
||||
|
||||
def test_fee_payout_pre_reset_or_lock_exists() -> None:
|
||||
"""FIX REQUIRED: pay-then-reset must become lock-then-pay-then-unlock.
|
||||
|
||||
wallet.py:1076-1080 currently: raw_send_to_lnurl() THEN reset_routstr_fee().
|
||||
A crash between these lines causes double payment.
|
||||
|
||||
Fix: set a lock flag BEFORE paying, clear it AFTER resetting.
|
||||
On startup, reconcile any locked-but-not-reset payouts.
|
||||
"""
|
||||
from routstr import wallet
|
||||
|
||||
source = inspect.getsource(wallet.periodic_routstr_fee_payout)
|
||||
|
||||
pay_pos = source.find("raw_send_to_lnurl")
|
||||
reset_pos = source.find("reset_routstr_fee")
|
||||
|
||||
assert pay_pos > 0 and reset_pos > 0, "Pay and reset both exist"
|
||||
|
||||
# After fix: lock/safeguard must exist BEFORE the pay call
|
||||
pre_pay_section = source[:pay_pos]
|
||||
has_pre_guard = any(
|
||||
kw in pre_pay_section.lower()
|
||||
for kw in ["lock", "payout_state", "is_paying", "in_progress",
|
||||
"pre_reset", "reconcile", "checkpoint"]
|
||||
)
|
||||
|
||||
assert has_pre_guard, (
|
||||
"FIX REQUIRED: Fee payout pays before resetting with no crash guard. "
|
||||
"A crash between pay and reset causes double payment. "
|
||||
"Fix: add a DB lock/payout_state flag before paying."
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# RED TESTS: DB store resilience
|
||||
# ===========================================================================
|
||||
|
||||
def test_retry_wrapper_exists() -> None:
|
||||
"""FIX REQUIRED: A retry wrapper for critical DB writes must exist."""
|
||||
from routstr.core import db
|
||||
|
||||
assert hasattr(db, "store_cashu_transaction_with_retry"), (
|
||||
"FIX REQUIRED: No retry wrapper exists for critical money-path "
|
||||
"DB writes. Was merged (#600) then reverted (#604). Must be "
|
||||
"reinstated with CRITICAL logging on final failure."
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Wallet caching mechanism (informational — not a bug on main)
|
||||
# ===========================================================================
|
||||
|
||||
def test_wallet_cache_uses_global_dict() -> None:
|
||||
"""get_wallet uses a global _wallets dict — verify mechanism."""
|
||||
from routstr import wallet
|
||||
|
||||
source = inspect.getsource(wallet.get_wallet)
|
||||
assert "_wallets" in source
|
||||
assert "load_mint" in source
|
||||
assert "load_proofs" in source
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Mint rate limiter setting (informational)
|
||||
# ===========================================================================
|
||||
|
||||
def test_mint_concurrency_setting_exists_or_documents_gap() -> None:
|
||||
"""If mint_max_concurrency exists, it must NOT be 0.
|
||||
|
||||
0 disables 429 cooldown tracking on the PR #597 branch.
|
||||
"""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
concurrency = getattr(settings, "mint_max_concurrency", None)
|
||||
if concurrency is not None:
|
||||
assert concurrency > 0, (
|
||||
f"mint_max_concurrency = {concurrency}. 0 disables 429 cooldown."
|
||||
)
|
||||
@@ -0,0 +1,185 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import AsyncGenerator
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
from sqlalchemy.pool import StaticPool
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.upstream.ehbp import (
|
||||
finalize_ehbp_actual_cost_payment,
|
||||
finalize_ehbp_max_cost_payment,
|
||||
)
|
||||
|
||||
|
||||
def _make_engine() -> AsyncEngine:
|
||||
return create_async_engine(
|
||||
"sqlite+aiosqlite://",
|
||||
poolclass=StaticPool,
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def session(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[AsyncSession, None]:
|
||||
monkeypatch.setattr("routstr.upstream.ehbp.ROUTSTR_FEE_PERCENT", 0)
|
||||
engine = _make_engine()
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
db_session = AsyncSession(engine, expire_on_commit=False)
|
||||
try:
|
||||
yield db_session
|
||||
finally:
|
||||
await db_session.close()
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def _api_key(session: AsyncSession, hashed_key: str) -> ApiKey | None:
|
||||
return (
|
||||
await session.exec(select(ApiKey).where(ApiKey.hashed_key == hashed_key))
|
||||
).one_or_none()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
key = ApiKey(
|
||||
hashed_key="ehbp-actual",
|
||||
balance=10_000,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
|
||||
await finalize_ehbp_actual_cost_payment(
|
||||
key,
|
||||
session,
|
||||
reserved_cost_for_model=3_000,
|
||||
model_id="tinfoil/model",
|
||||
cost_info={
|
||||
"total_msats": 1_200,
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 20,
|
||||
"input_msats": 500,
|
||||
"output_msats": 700,
|
||||
},
|
||||
)
|
||||
|
||||
updated = await _api_key(session, "ehbp-actual")
|
||||
assert updated is not None
|
||||
assert updated.balance == 8_800
|
||||
assert updated.reserved_balance == 0
|
||||
assert updated.reserved_at is None
|
||||
assert updated.total_spent == 1_200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_finalize_max_cost_payment_updates_parent_and_child_spend(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
parent = ApiKey(
|
||||
hashed_key="ehbp-parent",
|
||||
balance=10_000,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
)
|
||||
child = ApiKey(
|
||||
hashed_key="ehbp-child",
|
||||
balance=0,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
parent_key_hash="ehbp-parent",
|
||||
)
|
||||
session.add(parent)
|
||||
session.add(child)
|
||||
await session.commit()
|
||||
|
||||
await finalize_ehbp_max_cost_payment(
|
||||
child,
|
||||
session,
|
||||
max_cost_for_model=3_000,
|
||||
model_id="tinfoil/model",
|
||||
)
|
||||
|
||||
updated_parent = await _api_key(session, "ehbp-parent")
|
||||
updated_child = await _api_key(session, "ehbp-child")
|
||||
assert updated_parent is not None
|
||||
assert updated_child is not None
|
||||
assert updated_parent.balance == 7_000
|
||||
assert updated_parent.reserved_balance == 0
|
||||
assert updated_parent.reserved_at is None
|
||||
assert updated_parent.total_spent == 3_000
|
||||
assert updated_child.balance == 0
|
||||
assert updated_child.reserved_balance == 0
|
||||
assert updated_child.reserved_at is None
|
||||
assert updated_child.total_spent == 3_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matches_no_rows(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
key = ApiKey(
|
||||
hashed_key="ehbp-missing-parent",
|
||||
balance=10_000,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
await session.delete(key)
|
||||
await session.commit()
|
||||
|
||||
await finalize_ehbp_actual_cost_payment(
|
||||
key,
|
||||
session,
|
||||
reserved_cost_for_model=3_000,
|
||||
model_id="tinfoil/model",
|
||||
cost_info={"total_msats": 1_200},
|
||||
)
|
||||
|
||||
assert await _api_key(session, "ehbp-missing-parent") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_finalize_max_cost_payment_rolls_back_parent_when_child_update_matches_no_rows(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
parent = ApiKey(
|
||||
hashed_key="ehbp-rollback-parent",
|
||||
balance=10_000,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
)
|
||||
child = ApiKey(
|
||||
hashed_key="ehbp-missing-child",
|
||||
balance=0,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
parent_key_hash="ehbp-rollback-parent",
|
||||
)
|
||||
session.add(parent)
|
||||
session.add(child)
|
||||
await session.commit()
|
||||
await session.delete(child)
|
||||
await session.commit()
|
||||
|
||||
await finalize_ehbp_max_cost_payment(
|
||||
child,
|
||||
session,
|
||||
max_cost_for_model=3_000,
|
||||
model_id="tinfoil/model",
|
||||
)
|
||||
|
||||
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.reserved_at == 123
|
||||
assert updated_parent.total_spent == 0
|
||||
assert await _api_key(session, "ehbp-missing-child") is None
|
||||
@@ -0,0 +1,215 @@
|
||||
"""Tests asserting CORRECT behavior for emergency refund and DB persistence.
|
||||
|
||||
These tests FAIL against current main because the code is buggy.
|
||||
They serve as the "RED" phase of TDD — once the bugs are fixed, they go green.
|
||||
|
||||
Correct behavior required:
|
||||
1. store_cashu_transaction should raise on failure (not silently return False)
|
||||
2. Emergency refund paths must NOT use try/except/pass for DB stores
|
||||
3. A retry wrapper must exist for critical money-path DB writes
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
# ===========================================================================
|
||||
# RED TESTS: store_cashu_transaction should RAISE on failure
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_cashu_raises_on_db_failure_not_returns_false() -> None:
|
||||
"""FIX REQUIRED: store_cashu_transaction must raise on DB failure.
|
||||
|
||||
Currently returns False silently — callers never detect the failure.
|
||||
Correct behavior: raise an exception so callers can recover.
|
||||
"""
|
||||
from routstr.core.db import store_cashu_transaction
|
||||
|
||||
with patch("routstr.core.db.create_session") as mock_create:
|
||||
mock_session = AsyncMock()
|
||||
mock_session.commit = AsyncMock(side_effect=OSError("disk full"))
|
||||
mock_session.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_session.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_create.return_value = mock_session
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await store_cashu_transaction(
|
||||
token="cashuAtest_refund_token",
|
||||
amount=1000,
|
||||
unit="sat",
|
||||
mint_url="http://mint:3338",
|
||||
typ="out",
|
||||
request_id="req-123",
|
||||
)
|
||||
|
||||
# Must raise a meaningful exception, not silently return False
|
||||
# OSError or a custom DB error is acceptable
|
||||
assert "disk full" in str(exc_info.value) or isinstance(
|
||||
exc_info.value, (OSError, RuntimeError)
|
||||
), (
|
||||
f"Expected store to propagate the failure, got {type(exc_info.value).__name__}: "
|
||||
f"{exc_info.value}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_cashu_raises_on_any_error() -> None:
|
||||
"""FIX REQUIRED: All DB errors must propagate, not just OSError."""
|
||||
from routstr.core.db import store_cashu_transaction
|
||||
|
||||
errors = [
|
||||
OSError("disk full"),
|
||||
RuntimeError("connection lost"),
|
||||
ConnectionRefusedError("db down"),
|
||||
]
|
||||
|
||||
for error in errors:
|
||||
with patch("routstr.core.db.create_session") as mock_create:
|
||||
mock_session = AsyncMock()
|
||||
mock_session.commit = AsyncMock(side_effect=error)
|
||||
mock_session.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_session.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_create.return_value = mock_session
|
||||
|
||||
with pytest.raises(Exception):
|
||||
await store_cashu_transaction(
|
||||
token="cashuAtest",
|
||||
amount=1000,
|
||||
unit="sat",
|
||||
typ="out",
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# RED TESTS: Retry wrapper must exist
|
||||
# ===========================================================================
|
||||
|
||||
def test_retry_wrapper_exists_for_critical_writes() -> None:
|
||||
"""FIX REQUIRED: store_cashu_transaction_with_retry must exist.
|
||||
|
||||
Currently reverted (#600 → #604). All critical money-path DB writes
|
||||
(after minting a token) need retry with backoff + CRITICAL logging.
|
||||
"""
|
||||
from routstr.core import db
|
||||
|
||||
assert hasattr(db, "store_cashu_transaction_with_retry"), (
|
||||
"FIX REQUIRED: store_cashu_transaction_with_retry does not exist. "
|
||||
"Was merged in PR #600, reverted in PR #604. "
|
||||
"All post-mint DB writes need retry + backoff + CRITICAL logging."
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retry_wrapper_retries_on_transient_failure() -> None:
|
||||
"""FIX REQUIRED: retry wrapper must retry, not fail on first attempt."""
|
||||
from routstr.core import db
|
||||
|
||||
# Skip if the retry wrapper doesn't exist yet
|
||||
if not hasattr(db, "store_cashu_transaction_with_retry"):
|
||||
pytest.skip("store_cashu_transaction_with_retry does not exist yet")
|
||||
|
||||
with patch("routstr.core.db.store_cashu_transaction") as mock_store:
|
||||
mock_store = AsyncMock()
|
||||
mock_store.side_effect = [OSError("transient"), None] # 1st fails, 2nd succeeds
|
||||
# We'd test that the wrapper retries, but it doesn't exist yet
|
||||
# This test documents the expected behavior
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# RED TESTS: Emergency refund must not silently lose tokens
|
||||
# ===========================================================================
|
||||
|
||||
def test_emergency_refund_no_try_except_pass() -> None:
|
||||
"""FIX REQUIRED: Emergency refund paths must NOT use try/except/pass.
|
||||
|
||||
base.py:3643-3653 (chat) and base.py:4607-4617 (responses) both use
|
||||
try/except/pass around store_cashu_transaction after minting a refund
|
||||
token. If DB write fails, the token is permanently lost.
|
||||
|
||||
The fix: remove try/except/pass. Let the exception propagate so
|
||||
the caller can detect failure and at minimum log the token.
|
||||
"""
|
||||
import inspect
|
||||
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
# Check chat emergency refund handler
|
||||
chat_src = inspect.getsource(
|
||||
BaseUpstreamProvider.handle_x_cashu_non_streaming_response
|
||||
)
|
||||
|
||||
# Find the emergency refund section
|
||||
emergency_start = chat_src.find("emergency_refund = amount")
|
||||
assert emergency_start > 0, "Emergency refund path exists"
|
||||
|
||||
emergency_section = chat_src[emergency_start : emergency_start + 500]
|
||||
|
||||
# The try/except/pass around store_cashu_transaction must NOT exist
|
||||
has_except_pass = "except Exception:" in emergency_section and "pass" in emergency_section
|
||||
|
||||
assert not has_except_pass, (
|
||||
"FIX REQUIRED: Emergency refund (chat) uses try/except/pass around "
|
||||
"store_cashu_transaction. A failed DB write silently loses the minted "
|
||||
"token. Fix: let the exception propagate or log at CRITICAL with the "
|
||||
"full token for manual recovery."
|
||||
)
|
||||
|
||||
|
||||
def test_emergency_refund_responses_api_no_silent_failure() -> None:
|
||||
"""FIX REQUIRED: Responses API emergency refund same fix as chat."""
|
||||
import inspect
|
||||
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
responses_src = inspect.getsource(
|
||||
BaseUpstreamProvider.handle_x_cashu_non_streaming_responses_response
|
||||
)
|
||||
|
||||
has_emergency = "emergency_refund = amount" in responses_src
|
||||
if has_emergency:
|
||||
emergency_start = responses_src.find("emergency_refund = amount")
|
||||
emergency_section = responses_src[emergency_start : emergency_start + 500]
|
||||
has_except_pass = (
|
||||
"except Exception:" in emergency_section and "pass" in emergency_section
|
||||
)
|
||||
assert not has_except_pass, (
|
||||
"FIX REQUIRED: Responses API emergency refund also uses "
|
||||
"try/except/pass. Same fund-loss vulnerability as chat path."
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# RED TESTS: Fee payout crash safety
|
||||
# ===========================================================================
|
||||
|
||||
def test_fee_payout_has_crash_guard() -> None:
|
||||
"""FIX REQUIRED: Fee payout must have guard against double-pay on crash.
|
||||
|
||||
wallet.py:1076-1080 pays LNURL THEN resets the fee counter.
|
||||
A crash between these steps causes double payment on restart.
|
||||
|
||||
Fix options:
|
||||
1. Pre-reset the counter before paying (if pay fails, restore it)
|
||||
2. Add a "payout_lock" DB flag that's set before pay and cleared after
|
||||
3. Record payout in DB and reconcile on startup
|
||||
"""
|
||||
import inspect
|
||||
|
||||
from routstr import wallet
|
||||
|
||||
source = inspect.getsource(wallet.periodic_routstr_fee_payout)
|
||||
|
||||
# After the fix, the pay-then-reset pattern should be replaced
|
||||
# with a safe sequence. Verify the guard exists.
|
||||
has_guard = any(
|
||||
kw in source.lower()
|
||||
for kw in ["payout_lock", "is_paying", "payout_in_progress",
|
||||
"pre_reset", "reset_before", "reconcile"]
|
||||
)
|
||||
|
||||
assert has_guard, (
|
||||
"FIX REQUIRED: Fee payout has no crash guard. Pay-then-reset "
|
||||
"pattern in periodic_routstr_fee_payout can double-pay on "
|
||||
"process restart."
|
||||
)
|
||||
@@ -0,0 +1,160 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr import wallet
|
||||
from routstr.core import db
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _session_context(session: Mock) -> AsyncIterator[Mock]:
|
||||
yield session
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None:
|
||||
engine = create_async_engine("sqlite+aiosqlite://")
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
session.add(db.RoutstrFee(id=1, accumulated_msats=5_000))
|
||||
await session.commit()
|
||||
|
||||
assert await db.reset_routstr_fee(session, 5_000) is True
|
||||
assert await db.reset_routstr_fee(session, 5_000) is False
|
||||
|
||||
fee = await db.get_routstr_fee(session)
|
||||
await session.refresh(fee)
|
||||
assert fee.accumulated_msats == 0
|
||||
assert fee.payout_in_progress_msats == 5_000
|
||||
assert fee.total_paid_msats == 0
|
||||
|
||||
assert await db.complete_routstr_fee_payout(session, 5_000) is True
|
||||
await session.refresh(fee)
|
||||
assert fee.payout_in_progress_msats == 0
|
||||
assert fee.total_paid_msats == 5_000
|
||||
assert fee.last_paid_at is not None
|
||||
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_checkpoints_before_sending() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
payout_wallet = Mock()
|
||||
events: list[str] = []
|
||||
|
||||
async def checkpoint(*_args: object) -> bool:
|
||||
events.append("checkpoint")
|
||||
return True
|
||||
|
||||
async def send(*_args: object, **_kwargs: object) -> int:
|
||||
events.append("send")
|
||||
return 5
|
||||
|
||||
async def complete(*_args: object) -> bool:
|
||||
events.append("complete")
|
||||
return True
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", side_effect=checkpoint),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=payout_wallet)),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
assert events == ["checkpoint", "send", "complete"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_does_not_retry_an_unresolved_checkpoint() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=10_000,
|
||||
payout_in_progress_msats=5_000,
|
||||
payout_started_at=123,
|
||||
)
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock()) as checkpoint,
|
||||
patch("routstr.wallet.get_wallet", AsyncMock()) as get_wallet,
|
||||
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock()) as send,
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
checkpoint.assert_not_awaited()
|
||||
get_wallet.assert_not_awaited()
|
||||
send.assert_not_awaited()
|
||||
critical.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
complete = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", complete),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch(
|
||||
"routstr.wallet.raw_send_to_lnurl",
|
||||
AsyncMock(side_effect=TimeoutError("unknown outcome")),
|
||||
),
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
complete.assert_not_awaited()
|
||||
critical.assert_called_once()
|
||||
@@ -0,0 +1,46 @@
|
||||
import os
|
||||
import sqlite3
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _run_alembic(root: Path, database_url: str, revision: str) -> None:
|
||||
env = os.environ.copy()
|
||||
env["DATABASE_URL"] = database_url
|
||||
subprocess.run(
|
||||
[sys.executable, "-m", "alembic", "upgrade", revision],
|
||||
cwd=root,
|
||||
env=env,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
|
||||
def test_fee_payout_checkpoint_migration_preserves_existing_row(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
root = Path(__file__).resolve().parents[2]
|
||||
database_path = tmp_path / "migration.db"
|
||||
database_url = f"sqlite+aiosqlite:///{database_path}"
|
||||
_run_alembic(root, database_url, "c6d7e8f9a0b1")
|
||||
|
||||
with sqlite3.connect(database_path) as connection:
|
||||
result = connection.execute(
|
||||
"UPDATE routstr_fees SET accumulated_msats = 5000, "
|
||||
"total_paid_msats = 1000, last_paid_at = 123 WHERE id = 1"
|
||||
)
|
||||
assert result.rowcount == 1
|
||||
connection.commit()
|
||||
|
||||
_run_alembic(root, database_url, "head")
|
||||
|
||||
with sqlite3.connect(database_path) as connection:
|
||||
row = connection.execute(
|
||||
"SELECT accumulated_msats, total_paid_msats, last_paid_at, "
|
||||
"payout_in_progress_msats, payout_started_at "
|
||||
"FROM routstr_fees WHERE id = 1"
|
||||
).fetchone()
|
||||
|
||||
assert row == (5000, 1000, 123, 0, None)
|
||||
@@ -11,7 +11,9 @@ async def _fake_session(): # type: ignore[no-untyped-def]
|
||||
yield MagicMock()
|
||||
|
||||
|
||||
def _patches(proof_amount: int = 1000): # type: ignore[no-untyped-def]
|
||||
def _patches( # type: ignore[no-untyped-def]
|
||||
proof_amount: int = 1000, user_balance_msats: int = 0
|
||||
):
|
||||
proof = MagicMock(amount=proof_amount)
|
||||
return [
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
@@ -25,7 +27,7 @@ def _patches(proof_amount: int = 1000): # type: ignore[no-untyped-def]
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_for_mint_and_unit",
|
||||
AsyncMock(return_value=0),
|
||||
AsyncMock(return_value=user_balance_msats),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
]
|
||||
@@ -52,6 +54,31 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None:
|
||||
assert total_wallet == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_reports_liability_when_wallet_is_empty() -> None:
|
||||
"""An empty wallet must not hide outstanding user liabilities."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with patch.object(settings, "cashu_mints", []), patch.object(
|
||||
settings, "primary_mint", "http://primary:3338"
|
||||
):
|
||||
for p in _patches(proof_amount=0, user_balance_msats=5000):
|
||||
p.start()
|
||||
try:
|
||||
details, total_wallet, total_user, owner = await fetch_all_balances(
|
||||
units=["sat"]
|
||||
)
|
||||
finally:
|
||||
patch.stopall()
|
||||
|
||||
assert details[0]["wallet_balance"] == 0
|
||||
assert details[0]["user_balance"] == 5
|
||||
assert details[0]["owner_balance"] == -5
|
||||
assert total_wallet == 0
|
||||
assert total_user == 5
|
||||
assert owner == -5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_no_duplicate_primary_mint() -> None:
|
||||
"""primary_mint already in cashu_mints is not inspected twice."""
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
"""raw_send_to_lnurl() must not hang forever on an unresponsive mint.
|
||||
|
||||
The Cashu library issues POST /v1/melt/bolt11 with timeout=None, so a hung
|
||||
mint would block the melt (and the payout loop) indefinitely. raw_send_to_lnurl
|
||||
now wraps wallet.melt() in asyncio.wait_for(MELT_TIMEOUT_SECONDS) and surfaces a
|
||||
timeout as LNURLError instead of hanging.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.payment import lnurl
|
||||
from routstr.payment.lnurl import LNURLError, raw_send_to_lnurl
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_times_out_on_hung_melt() -> None:
|
||||
proofs = [MagicMock(amount=1000)]
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q"))
|
||||
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
|
||||
|
||||
async def _hang(**kwargs: object) -> None:
|
||||
await asyncio.sleep(5) # far longer than the patched timeout
|
||||
|
||||
wallet.melt = AsyncMock(side_effect=_hang)
|
||||
|
||||
lnurl_data = {
|
||||
"callback_url": "https://ln.tld/cb",
|
||||
"min_sendable": 1_000,
|
||||
"max_sendable": 100_000_000,
|
||||
}
|
||||
|
||||
with patch.object(lnurl, "MELT_TIMEOUT_SECONDS", 0.05), patch(
|
||||
"routstr.payment.lnurl.get_lnurl_data", AsyncMock(return_value=lnurl_data)
|
||||
), patch(
|
||||
"routstr.payment.lnurl.get_lnurl_invoice",
|
||||
AsyncMock(return_value=("lnbc1...", {})),
|
||||
):
|
||||
with pytest.raises(LNURLError, match="Melt timed out"):
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_succeeds_within_timeout() -> None:
|
||||
"""A prompt melt still returns the net amount, unaffected by the guard."""
|
||||
proofs = [MagicMock(amount=1000)]
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q"))
|
||||
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
|
||||
wallet.melt = AsyncMock(return_value=MagicMock())
|
||||
|
||||
lnurl_data = {
|
||||
"callback_url": "https://ln.tld/cb",
|
||||
"min_sendable": 1_000,
|
||||
"max_sendable": 100_000_000,
|
||||
}
|
||||
|
||||
with patch.object(lnurl, "MELT_TIMEOUT_SECONDS", 5), patch(
|
||||
"routstr.payment.lnurl.get_lnurl_data", AsyncMock(return_value=lnurl_data)
|
||||
), patch(
|
||||
"routstr.payment.lnurl.get_lnurl_invoice",
|
||||
AsyncMock(return_value=("lnbc1...", {})),
|
||||
):
|
||||
paid = await raw_send_to_lnurl(
|
||||
wallet, proofs, "owner@ln.tld", "sat", amount=1000
|
||||
)
|
||||
|
||||
assert paid > 0
|
||||
wallet.melt.assert_awaited_once()
|
||||
@@ -502,7 +502,12 @@ async def test_streaming_emits_sse_and_reconciles_cost_at_end() -> None:
|
||||
captured_cost_call: dict[str, Any] = {}
|
||||
|
||||
async def fake_adjust(
|
||||
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
|
||||
fresh_key: Any,
|
||||
combined_data: Any,
|
||||
sess: Any,
|
||||
max_cost: int,
|
||||
model_obj: Any = None,
|
||||
provider_fee: Any = None,
|
||||
) -> dict:
|
||||
captured_cost_call["combined_data"] = combined_data
|
||||
captured_cost_call["max_cost"] = max_cost
|
||||
@@ -591,7 +596,12 @@ async def test_streaming_handles_iterator_yielding_raw_sse_bytes() -> None:
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
async def fake_adjust(
|
||||
fresh_key: Any, combined_data: Any, sess: Any, max_cost: int, usage: Any = None
|
||||
fresh_key: Any,
|
||||
combined_data: Any,
|
||||
sess: Any,
|
||||
max_cost: int,
|
||||
model_obj: Any = None,
|
||||
provider_fee: Any = None,
|
||||
) -> dict:
|
||||
captured["combined_data"] = combined_data
|
||||
return fake_cost
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
"""Tests for periodic_payout() resilience fixes.
|
||||
|
||||
Covers two regressions from the auto-payout / primary-mint audit
|
||||
(docs/auto-payout-primary-mint-failure-report.md):
|
||||
|
||||
1. periodic_payout() must include settings.primary_mint even when it is not
|
||||
listed in settings.cashu_mints, matching fetch_all_balances(); otherwise
|
||||
primary-mint funds never auto-payout.
|
||||
2. A failure on one mint/unit must not abort payout for the remaining
|
||||
mint/units in the same cycle (the try/except is now per mint/unit).
|
||||
"""
|
||||
|
||||
from collections.abc import Callable, Coroutine
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.wallet import periodic_payout
|
||||
|
||||
# Sentinel interval used to break the otherwise-infinite payout loop after
|
||||
# exactly one full cycle.
|
||||
_INTERVAL = 987
|
||||
|
||||
|
||||
class _LoopBreak(Exception):
|
||||
"""Raised via the patched sleep to stop periodic_payout after one cycle."""
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _fake_session(): # type: ignore[no-untyped-def]
|
||||
yield MagicMock()
|
||||
|
||||
|
||||
def _one_cycle_sleep() -> Callable[[float], Coroutine[Any, Any, None]]:
|
||||
"""Return an async sleep stub that lets exactly one payout cycle run.
|
||||
|
||||
The top-of-loop sleep uses the sentinel interval; the second time it is
|
||||
seen (start of the second cycle) we raise to break out. The inner
|
||||
``asyncio.sleep(5)`` pass-through is ignored.
|
||||
"""
|
||||
seen = {"interval": 0}
|
||||
|
||||
async def _sleep(seconds: float) -> None:
|
||||
if seconds == _INTERVAL:
|
||||
seen["interval"] += 1
|
||||
if seen["interval"] >= 2:
|
||||
raise _LoopBreak()
|
||||
|
||||
return _sleep
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> None:
|
||||
"""primary_mint absent from cashu_mints is still paid out."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
get_wallet = AsyncMock(return_value=MagicMock())
|
||||
raw_send = AsyncMock(return_value=1000)
|
||||
|
||||
with patch.object(settings, "cashu_mints", []), patch.object(
|
||||
settings, "primary_mint", "http://primary:3338"
|
||||
), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object(
|
||||
settings, "payout_interval_seconds", _INTERVAL
|
||||
), patch.object(settings, "min_payout_sat", 10), patch(
|
||||
"routstr.wallet.asyncio.sleep", _one_cycle_sleep()
|
||||
), patch("routstr.wallet.db.create_session", _fake_session), patch(
|
||||
"routstr.wallet.get_wallet", get_wallet
|
||||
), patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
), patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
), patch(
|
||||
"routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0)
|
||||
), patch("routstr.wallet.raw_send_to_lnurl", raw_send):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
processed = {call.args[0] for call in get_wallet.await_args_list}
|
||||
assert processed == {"http://primary:3338"}
|
||||
assert raw_send.await_count >= 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_periodic_payout_isolates_failing_mint() -> None:
|
||||
"""A failing mint does not prevent payout for the other mints."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
async def _get_wallet(mint_url: str, unit: str) -> MagicMock:
|
||||
if mint_url == "http://bad:3338":
|
||||
raise RuntimeError("mint unreachable")
|
||||
return MagicMock()
|
||||
|
||||
get_wallet = AsyncMock(side_effect=_get_wallet)
|
||||
raw_send = AsyncMock(return_value=1000)
|
||||
|
||||
with patch.object(
|
||||
settings, "cashu_mints", ["http://bad:3338", "http://good:3338"]
|
||||
), patch.object(settings, "primary_mint", "http://good:3338"), patch.object(
|
||||
settings, "receive_ln_address", "owner@ln.tld"
|
||||
), patch.object(settings, "payout_interval_seconds", _INTERVAL), patch.object(
|
||||
settings, "min_payout_sat", 10
|
||||
), patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), patch(
|
||||
"routstr.wallet.db.create_session", _fake_session
|
||||
), patch("routstr.wallet.get_wallet", get_wallet), patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
), patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
), patch(
|
||||
"routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0)
|
||||
), patch("routstr.wallet.raw_send_to_lnurl", raw_send):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
# The bad mint raised on get_wallet for both units, yet the good mint was
|
||||
# still reached and paid out for both units — failures are isolated.
|
||||
good_calls = [
|
||||
c for c in get_wallet.await_args_list if c.args[0] == "http://good:3338"
|
||||
]
|
||||
assert len(good_calls) == 2 # sat + msat
|
||||
assert raw_send.await_count == 2 # good mint paid for both units
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_periodic_payout_handles_session_creation_failure() -> None:
|
||||
"""A db.create_session failure is logged and the payout loop continues."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
create_session = MagicMock(side_effect=RuntimeError("db unavailable"))
|
||||
logger = MagicMock()
|
||||
|
||||
with patch.object(settings, "cashu_mints", ["http://mint:3338"]), patch.object(
|
||||
settings, "primary_mint", "http://mint:3338"
|
||||
), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object(
|
||||
settings, "payout_interval_seconds", _INTERVAL
|
||||
), patch(
|
||||
"routstr.wallet.asyncio.sleep", _one_cycle_sleep()
|
||||
), patch(
|
||||
"routstr.wallet.db.create_session", create_session
|
||||
), patch("routstr.wallet.logger", logger):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
create_session.assert_called_once()
|
||||
logger.error.assert_called_once()
|
||||
message = logger.error.call_args.args[0]
|
||||
extra = logger.error.call_args.kwargs["extra"]
|
||||
assert message == "Error in periodic payout cycle: RuntimeError"
|
||||
assert extra == {"error": "db unavailable"}
|
||||
@@ -0,0 +1,126 @@
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import AsyncGenerator
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
from sqlalchemy.pool import StaticPool
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr import wallet
|
||||
from routstr.core.db import CashuTransaction
|
||||
|
||||
|
||||
def _make_engine() -> AsyncEngine:
|
||||
return create_async_engine(
|
||||
"sqlite+aiosqlite://",
|
||||
poolclass=StaticPool,
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def refund_db(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> AsyncGenerator[AsyncEngine, None]:
|
||||
engine = _make_engine()
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
@asynccontextmanager
|
||||
async def create_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
yield session
|
||||
|
||||
monkeypatch.setattr(wallet.db, "create_session", create_session)
|
||||
try:
|
||||
yield engine
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def _store(
|
||||
engine: AsyncEngine,
|
||||
*,
|
||||
token: str,
|
||||
created_at: int,
|
||||
typ: str = "out",
|
||||
collected: bool = False,
|
||||
swept: bool = False,
|
||||
) -> None:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
session.add(
|
||||
CashuTransaction(
|
||||
token=token,
|
||||
amount=10,
|
||||
unit="sat",
|
||||
type=typ,
|
||||
created_at=created_at,
|
||||
collected=collected,
|
||||
swept=swept,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def _transactions(engine: AsyncEngine) -> dict[str, CashuTransaction]:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
results = await session.exec(select(CashuTransaction))
|
||||
return {transaction.token: transaction for transaction in results.all()}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_only_processes_eligible_outbound_transactions(
|
||||
refund_db: AsyncEngine, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cutoff = 1_000
|
||||
await _store(refund_db, token="eligible", created_at=cutoff - 1)
|
||||
await _store(refund_db, token="too-new", created_at=cutoff)
|
||||
await _store(
|
||||
refund_db, token="collected", created_at=cutoff - 1, collected=True
|
||||
)
|
||||
await _store(refund_db, token="swept", created_at=cutoff - 1, swept=True)
|
||||
await _store(refund_db, token="inbound", created_at=cutoff - 1, typ="in")
|
||||
receive_token = AsyncMock()
|
||||
monkeypatch.setattr(wallet, "recieve_token", receive_token)
|
||||
|
||||
await wallet._refund_sweep_once(cutoff)
|
||||
|
||||
receive_token.assert_awaited_once_with("eligible")
|
||||
transactions = await _transactions(refund_db)
|
||||
assert transactions["eligible"].swept is True
|
||||
assert transactions["too-new"].swept is False
|
||||
assert transactions["collected"].collected is True
|
||||
assert transactions["collected"].swept is False
|
||||
assert transactions["swept"].swept is True
|
||||
assert transactions["inbound"].swept is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_persists_success_and_isolates_token_failures(
|
||||
refund_db: AsyncEngine, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
cutoff = 1_000
|
||||
for token in ("success", "already-spent", "temporary-failure"):
|
||||
await _store(refund_db, token=token, created_at=cutoff - 1)
|
||||
|
||||
async def receive_token(token: str) -> None:
|
||||
if token == "already-spent":
|
||||
raise RuntimeError("Token already spent")
|
||||
if token == "temporary-failure":
|
||||
raise RuntimeError("mint unavailable")
|
||||
|
||||
receive = AsyncMock(side_effect=receive_token)
|
||||
monkeypatch.setattr(wallet, "recieve_token", receive)
|
||||
|
||||
await wallet._refund_sweep_once(cutoff)
|
||||
|
||||
assert receive.await_count == 3
|
||||
transactions = await _transactions(refund_db)
|
||||
assert transactions["success"].swept is True
|
||||
assert transactions["success"].collected is False
|
||||
assert transactions["already-spent"].collected is True
|
||||
assert transactions["already-spent"].swept is False
|
||||
assert transactions["temporary-failure"].collected is False
|
||||
assert transactions["temporary-failure"].swept is False
|
||||
@@ -0,0 +1,141 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.responses import Response
|
||||
from httpx import ASGITransport, AsyncClient
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def proxy_app() -> FastAPI:
|
||||
app = FastAPI()
|
||||
app.include_router(proxy_module.proxy_router)
|
||||
return app
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attestation_get_routes_directly_to_tinfoil_provider(
|
||||
monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI
|
||||
) -> None:
|
||||
non_tinfoil = MagicMock()
|
||||
non_tinfoil.provider_type = "openai"
|
||||
non_tinfoil.prepare_headers = MagicMock(return_value={})
|
||||
non_tinfoil.forward_get_request = AsyncMock(
|
||||
return_value=Response(status_code=404, content=b"wrong upstream")
|
||||
)
|
||||
|
||||
tinfoil = MagicMock()
|
||||
tinfoil.provider_type = "tinfoil"
|
||||
tinfoil.prepare_headers = MagicMock(return_value={"accept": "application/json"})
|
||||
tinfoil.forward_get_request = AsyncMock(
|
||||
return_value=Response(status_code=200, content=b'{"attestation":true}')
|
||||
)
|
||||
|
||||
monkeypatch.setattr(proxy_module, "_upstreams", [non_tinfoil, tinfoil])
|
||||
|
||||
async with AsyncClient(
|
||||
transport=ASGITransport(app=proxy_app), # type: ignore[arg-type]
|
||||
base_url="http://test",
|
||||
) as client:
|
||||
response = await client.get("/attestation")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.content == b'{"attestation":true}'
|
||||
non_tinfoil.forward_get_request.assert_not_called()
|
||||
tinfoil.forward_get_request.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tee_attestation_get_routes_directly_to_tinfoil_provider(
|
||||
monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI
|
||||
) -> None:
|
||||
non_tinfoil = MagicMock()
|
||||
non_tinfoil.provider_type = "openrouter"
|
||||
non_tinfoil.prepare_headers = MagicMock(return_value={})
|
||||
non_tinfoil.forward_get_request = AsyncMock(
|
||||
return_value=Response(status_code=404, content=b"wrong upstream")
|
||||
)
|
||||
|
||||
tinfoil = MagicMock()
|
||||
tinfoil.provider_type = "tinfoil"
|
||||
tinfoil.prepare_headers = MagicMock(return_value={"accept": "application/json"})
|
||||
tinfoil.forward_get_request = AsyncMock(
|
||||
return_value=Response(status_code=200, content=b'{"tee":true}')
|
||||
)
|
||||
|
||||
monkeypatch.setattr(proxy_module, "_upstreams", [non_tinfoil, tinfoil])
|
||||
|
||||
async with AsyncClient(
|
||||
transport=ASGITransport(app=proxy_app), # type: ignore[arg-type]
|
||||
base_url="http://test",
|
||||
) as client:
|
||||
response = await client.get("/tee/attestation")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.content == b'{"tee":true}'
|
||||
non_tinfoil.forward_get_request.assert_not_called()
|
||||
tinfoil.forward_get_request.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ["attestation/", "tee/attestation/"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_attestation_trailing_slash_routes_directly_to_tinfoil(
|
||||
monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI, path: str
|
||||
) -> None:
|
||||
tinfoil = MagicMock()
|
||||
tinfoil.provider_type = "tinfoil"
|
||||
tinfoil.prepare_headers = MagicMock(return_value={})
|
||||
tinfoil.forward_get_request = AsyncMock(return_value=Response(status_code=200))
|
||||
monkeypatch.setattr(proxy_module, "_upstreams", [tinfoil])
|
||||
|
||||
async with AsyncClient(
|
||||
transport=ASGITransport(app=proxy_app), # type: ignore[arg-type]
|
||||
base_url="http://test",
|
||||
) as client:
|
||||
response = await client.get(f"/{path}")
|
||||
|
||||
assert response.status_code == 200
|
||||
tinfoil.forward_get_request.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ["attestation/foo", "attestationjunk"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_attestation_prefix_does_not_bypass_authentication(
|
||||
monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI, path: str
|
||||
) -> None:
|
||||
tinfoil = MagicMock()
|
||||
tinfoil.provider_type = "tinfoil"
|
||||
tinfoil.forward_get_request = AsyncMock()
|
||||
monkeypatch.setattr(proxy_module, "_upstreams", [tinfoil])
|
||||
|
||||
async with AsyncClient(
|
||||
transport=ASGITransport(app=proxy_app), # type: ignore[arg-type]
|
||||
base_url="http://test",
|
||||
) as client:
|
||||
response = await client.get(f"/{path}")
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.json()["error"]["type"] == "invalid_model"
|
||||
tinfoil.forward_get_request.assert_not_awaited()
|
||||
|
||||
|
||||
def test_attestation_upstream_selection_is_tinfoil_only() -> None:
|
||||
non_tinfoil = MagicMock(provider_type="openai")
|
||||
tinfoil = MagicMock(provider_type="tinfoil")
|
||||
|
||||
assert proxy_module._select_unauthenticated_get_upstreams(
|
||||
"attestation", [non_tinfoil, tinfoil]
|
||||
) == [tinfoil]
|
||||
assert proxy_module._select_unauthenticated_get_upstreams(
|
||||
"tee/attestation", [non_tinfoil, tinfoil]
|
||||
) == [tinfoil]
|
||||
assert proxy_module._select_unauthenticated_get_upstreams(
|
||||
"attestation/", [non_tinfoil, tinfoil]
|
||||
) == [tinfoil]
|
||||
assert proxy_module._select_unauthenticated_get_upstreams(
|
||||
"attestationjunk", [non_tinfoil, tinfoil]
|
||||
) == [non_tinfoil, tinfoil]
|
||||
@@ -0,0 +1,118 @@
|
||||
from collections.abc import AsyncIterator
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import CashuTransaction
|
||||
from routstr.wallet import refund_sweep_once
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def session_factory(
|
||||
tmp_path: Path,
|
||||
) -> AsyncIterator[async_sessionmaker[AsyncSession]]:
|
||||
engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'refunds.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()
|
||||
|
||||
|
||||
async def _insert(
|
||||
factory: async_sessionmaker[AsyncSession], *rows: CashuTransaction
|
||||
) -> None:
|
||||
async with factory() as session:
|
||||
session.add_all(rows)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def _load(
|
||||
factory: async_sessionmaker[AsyncSession],
|
||||
) -> dict[str, CashuTransaction]:
|
||||
async with factory() as session:
|
||||
result = await session.exec(select(CashuTransaction))
|
||||
return {row.token: row for row in result.all()}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_only_processes_expired_eligible_outgoing_tokens(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
rows = [
|
||||
CashuTransaction(
|
||||
token="eligible", amount=1, unit="sat", type="out", created_at=800
|
||||
),
|
||||
CashuTransaction(
|
||||
token="fresh", amount=1, unit="sat", type="out", created_at=950
|
||||
),
|
||||
CashuTransaction(
|
||||
token="boundary", amount=1, unit="sat", type="out", created_at=900
|
||||
),
|
||||
CashuTransaction(
|
||||
token="collected",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
collected=True,
|
||||
),
|
||||
CashuTransaction(
|
||||
token="swept", amount=1, unit="sat", type="out", created_at=800, swept=True
|
||||
),
|
||||
CashuTransaction(
|
||||
token="incoming", amount=1, unit="sat", type="in", created_at=800
|
||||
),
|
||||
]
|
||||
await _insert(session_factory, *rows)
|
||||
receive = AsyncMock()
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch("routstr.wallet.recieve_token", receive),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
receive.assert_awaited_once_with("eligible")
|
||||
loaded = await _load(session_factory)
|
||||
assert loaded["eligible"].swept is True
|
||||
assert loaded["fresh"].swept is False
|
||||
assert loaded["boundary"].swept is False
|
||||
assert loaded["collected"].swept is False
|
||||
assert loaded["swept"].swept is True
|
||||
assert loaded["incoming"].swept is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("error", "collected"),
|
||||
[
|
||||
(RuntimeError("token already spent"), True),
|
||||
(RuntimeError("mint unavailable"), False),
|
||||
],
|
||||
)
|
||||
async def test_refund_sweep_records_terminal_but_not_transient_failures(
|
||||
session_factory: async_sessionmaker[AsyncSession], error: Exception, collected: bool
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="refund", amount=1, unit="sat", type="out", created_at=800
|
||||
),
|
||||
)
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch("routstr.wallet.recieve_token", AsyncMock(side_effect=error)),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
refund = (await _load(session_factory))["refund"]
|
||||
assert refund.collected is collected
|
||||
assert refund.swept is False
|
||||
+132
-1
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
@@ -6,7 +7,15 @@ from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import text
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.settings import Settings, SettingsService
|
||||
from routstr.core.settings import Settings, SettingsService, settings
|
||||
|
||||
NSEC_HEX = "1" * 64
|
||||
|
||||
|
||||
async def _read_settings_blob(session: AsyncSession) -> dict:
|
||||
"""Return the raw persisted settings JSON (id=1) as a dict."""
|
||||
row = await session.exec(text("SELECT data FROM settings WHERE id = 1")) # type: ignore
|
||||
return json.loads(row.first()[0])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -120,3 +129,125 @@ async def test_settings_initialize_discards_unknown_keys() -> None:
|
||||
assert '"enable_analytics_sharing": true' in stored_data
|
||||
assert "nostr_analytics_enabled" not in stored_data
|
||||
assert "unknown_key" not in stored_data
|
||||
|
||||
|
||||
# ── Secret fields are never written to the settings blob (issue #553) ────────
|
||||
|
||||
|
||||
def test_settings_model_drops_admin_password_field() -> None:
|
||||
# admin_password now lives only as a one-way hash in the Secret store; it is
|
||||
# no longer a settings field at all.
|
||||
assert "admin_password" not in Settings.__fields__
|
||||
# nsec remains a runtime value held in memory; upstream_api_key is ordinary
|
||||
# config that still lives in the persisted blob.
|
||||
assert "nsec" in Settings.__fields__
|
||||
assert "upstream_api_key" in Settings.__fields__
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_secret_fields_kept_in_memory_but_not_persisted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# bootstrap_secrets (which runs first at boot) owns the nsec and has already
|
||||
# decrypted it into memory; simulate that live value. initialize must keep it
|
||||
# in memory for runtime consumers yet never write it to the settings blob.
|
||||
monkeypatch.setattr(settings, "nsec", NSEC_HEX)
|
||||
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
s = await SettingsService.initialize(session)
|
||||
|
||||
# Runtime consumers still see the live secret value.
|
||||
assert s.nsec == NSEC_HEX
|
||||
|
||||
# ...but it is never written to the settings blob.
|
||||
blob = await _read_settings_blob(session)
|
||||
assert "nsec" not in blob
|
||||
assert "admin_password" not in blob
|
||||
# Non-secret derived/public values are still persisted.
|
||||
assert blob["npub"] == s.npub
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_api_key_survives_persistence(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# upstream_api_key is provider-scoped config, not a vault secret: it has no
|
||||
# encrypted home yet, so it must stay in the settings blob. Stripping it
|
||||
# would load it once, rewrite the blob without it, and lose it on the next
|
||||
# restart. Guard the on-disk survival path: blob-only value, no env.
|
||||
monkeypatch.delenv("UPSTREAM_API_KEY", raising=False)
|
||||
monkeypatch.setattr(settings, "upstream_api_key", "")
|
||||
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
await SettingsService.initialize(session)
|
||||
await session.exec( # type: ignore
|
||||
text("UPDATE settings SET data = :d WHERE id = 1").bindparams(
|
||||
d=json.dumps({"name": "LegacyNode", "upstream_api_key": "sk-only-in-db"})
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
# A reload must not drop the key from the blob...
|
||||
await SettingsService.initialize(session)
|
||||
blob = await _read_settings_blob(session)
|
||||
assert blob["upstream_api_key"] == "sk-only-in-db"
|
||||
# ...and it stays live for the proxy hot path.
|
||||
assert settings.upstream_api_key == "sk-only-in-db"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_existing_blob_secrets_are_stripped_on_initialize(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.delenv("NSEC", raising=False)
|
||||
monkeypatch.delenv("UPSTREAM_API_KEY", raising=False)
|
||||
monkeypatch.delenv("ADMIN_PASSWORD", raising=False)
|
||||
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
await SettingsService.initialize(session)
|
||||
|
||||
# Simulate a legacy row that still carries plaintext secrets in the blob.
|
||||
await session.exec( # type: ignore
|
||||
text("UPDATE settings SET data = :d WHERE id = 1").bindparams(
|
||||
d=json.dumps(
|
||||
{
|
||||
"name": "LegacyNode",
|
||||
"admin_password": "pw",
|
||||
"nsec": NSEC_HEX,
|
||||
"upstream_api_key": "sk-legacy",
|
||||
}
|
||||
)
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
await SettingsService.initialize(session)
|
||||
|
||||
blob = await _read_settings_blob(session)
|
||||
assert "admin_password" not in blob
|
||||
assert "nsec" not in blob
|
||||
# Non-secret values survive the migration, including upstream_api_key,
|
||||
# which is not vaulted yet and so must stay in the blob.
|
||||
assert blob["name"] == "LegacyNode"
|
||||
assert blob["upstream_api_key"] == "sk-legacy"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_does_not_persist_secret_fields(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("NSEC", NSEC_HEX)
|
||||
monkeypatch.setattr(settings, "nsec", "")
|
||||
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
await SettingsService.initialize(session)
|
||||
await SettingsService.update({"name": "Updated"}, session)
|
||||
|
||||
blob = await _read_settings_blob(session)
|
||||
assert blob["name"] == "Updated"
|
||||
assert "nsec" not in blob
|
||||
assert "admin_password" not in blob
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
"""Settlement bills the model and provider fee that actually served.
|
||||
|
||||
Covers ``calculate_cost``'s served-identity parameters: a passed ``model_obj``
|
||||
is billed directly instead of re-deriving pricing from the response's model
|
||||
string through the alias map (which yields the best-ranked candidate, not the
|
||||
serving one), and a passed ``provider_fee`` is applied on the USD-cost path
|
||||
instead of the best-ranked provider's fee. The string/alias fallbacks remain
|
||||
for callers without routed identity.
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||
|
||||
from routstr.payment.cost_calculation import CostData, calculate_cost
|
||||
from routstr.payment.models import Architecture, Model, Pricing
|
||||
|
||||
|
||||
def _make_model(
|
||||
model_id: str, prompt_sats: float, completion_sats: float
|
||||
) -> Model:
|
||||
return Model(
|
||||
id=model_id,
|
||||
name=model_id,
|
||||
created=0,
|
||||
description="",
|
||||
context_length=64000,
|
||||
architecture=Architecture(
|
||||
modality="text->text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="Other",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(prompt=prompt_sats, completion=completion_sats),
|
||||
sats_pricing=Pricing(prompt=prompt_sats, completion=completion_sats),
|
||||
)
|
||||
|
||||
|
||||
WINNER = _make_model("dual-model", 0.001, 0.002)
|
||||
SERVED = _make_model("dual-model", 0.005, 0.010)
|
||||
|
||||
RESPONSE = {
|
||||
"model": "dual-model",
|
||||
"usage": {
|
||||
"prompt_tokens": 1000,
|
||||
"completion_tokens": 500,
|
||||
"total_tokens": 1500,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_sats_usd_price() -> None: # type: ignore[misc]
|
||||
with patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-4
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_served_model_pricing_wins_over_alias_lookup() -> None:
|
||||
"""With ``model_obj`` given, the alias map is not consulted for pricing."""
|
||||
with patch(
|
||||
"routstr.proxy.get_model_instance", return_value=WINNER
|
||||
) as alias_lookup:
|
||||
result = await calculate_cost(
|
||||
dict(RESPONSE), max_cost=100_000, model_obj=SERVED
|
||||
)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
# 1000/1000 * 5000 + 500/1000 * 10000 msats at the SERVED model's rates.
|
||||
assert result.total_msats == 10_000
|
||||
alias_lookup.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_string_fallback_still_prices_without_model_obj() -> None:
|
||||
"""Callers without routed identity keep the alias-map string lookup."""
|
||||
with patch("routstr.proxy.get_model_instance", return_value=WINNER):
|
||||
result = await calculate_cost(dict(RESPONSE), max_cost=100_000)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.total_msats == 2_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usd_cost_path_applies_given_provider_fee() -> None:
|
||||
"""The USD-cost path bills the serving provider's fee when supplied."""
|
||||
from unittest.mock import Mock
|
||||
|
||||
response = dict(RESPONSE)
|
||||
response["usage"] = dict(RESPONSE["usage"], cost=0.001) # type: ignore[arg-type]
|
||||
|
||||
best_ranked = Mock(provider_fee=1.0)
|
||||
with patch(
|
||||
"routstr.proxy.get_provider_for_model", return_value=[best_ranked]
|
||||
):
|
||||
result = await calculate_cost(
|
||||
response, max_cost=100_000, model_obj=SERVED, provider_fee=1.5
|
||||
)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
# 0.001 USD * fee 1.5 / 0.0005 USD-per-sat = 3 sats = 3000 msats.
|
||||
assert result.total_msats == 3_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usd_cost_path_falls_back_to_best_ranked_fee() -> None:
|
||||
"""Without a supplied fee, the alias-map provider lookup still applies."""
|
||||
from unittest.mock import Mock
|
||||
|
||||
response = dict(RESPONSE)
|
||||
response["usage"] = dict(RESPONSE["usage"], cost=0.001) # type: ignore[arg-type]
|
||||
|
||||
best_ranked = Mock(provider_fee=2.0)
|
||||
with patch(
|
||||
"routstr.proxy.get_provider_for_model", return_value=[best_ranked]
|
||||
):
|
||||
result = await calculate_cost(response, max_cost=100_000)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.total_msats == 4_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_x_cashu_cost_uses_served_model_not_upstream_echo() -> None:
|
||||
"""``get_x_cashu_cost`` bills the routed model, not the raw model echo.
|
||||
|
||||
X-Cashu handlers do not rewrite the upstream's echoed model string, so
|
||||
without the routed model the settle would look up whatever wire name the
|
||||
upstream reported. With ``model_obj`` given, the echo must be irrelevant.
|
||||
"""
|
||||
from routstr.upstream import GenericUpstreamProvider
|
||||
|
||||
provider = GenericUpstreamProvider("http://upstream.example", "key", 1.0)
|
||||
response = dict(RESPONSE, model="totally-unknown-wire-name")
|
||||
|
||||
with patch(
|
||||
"routstr.proxy.get_model_instance", return_value=WINNER
|
||||
) as alias_lookup:
|
||||
cost = await provider.get_x_cashu_cost(
|
||||
response, max_cost_for_model=100_000, model_obj=SERVED
|
||||
)
|
||||
|
||||
assert cost is not None
|
||||
assert cost.total_msats == 10_000
|
||||
alias_lookup.assert_not_called()
|
||||
@@ -355,8 +355,11 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
|
||||
revert_mock = AsyncMock(return_value=True)
|
||||
|
||||
with (
|
||||
patch.object(proxy_module, "get_model_instance", return_value=MagicMock()),
|
||||
patch.object(proxy_module, "get_provider_for_model", return_value=[upstream]),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"get_candidates",
|
||||
return_value=[(MagicMock(), upstream)],
|
||||
),
|
||||
patch.object(
|
||||
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
|
||||
),
|
||||
|
||||
@@ -0,0 +1,749 @@
|
||||
"""Unit tests for Tinfoil direct integration.
|
||||
|
||||
Covers the EHBP usage-metrics header parser, the proxy header stripping, the
|
||||
enclave URL override, and the TinfoilUpstreamProvider model fetching/forwarding
|
||||
target logic.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.upstream.ehbp import (
|
||||
_PROXY_ONLY_HEADERS,
|
||||
_compute_ehbp_actual_cost,
|
||||
_prepare_ehbp_upstream_headers,
|
||||
_resolve_ehbp_target_url,
|
||||
_strip_proxy_headers,
|
||||
parse_tinfoil_usage_metrics,
|
||||
)
|
||||
from routstr.upstream.tinfoil import (
|
||||
TinfoilModel,
|
||||
TinfoilUpstreamProvider,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# parse_tinfoil_usage_metrics
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestParseTinfoilUsageMetrics:
|
||||
def test_full_header(self) -> None:
|
||||
result = parse_tinfoil_usage_metrics("prompt=67,completion=42,total=109")
|
||||
assert result == {
|
||||
"prompt_tokens": 67,
|
||||
"completion_tokens": 42,
|
||||
"total_tokens": 109,
|
||||
}
|
||||
|
||||
def test_without_total(self) -> None:
|
||||
result = parse_tinfoil_usage_metrics("prompt=10,completion=5")
|
||||
assert result == {"prompt_tokens": 10, "completion_tokens": 5}
|
||||
|
||||
def test_none(self) -> None:
|
||||
assert parse_tinfoil_usage_metrics(None) is None
|
||||
|
||||
def test_empty(self) -> None:
|
||||
assert parse_tinfoil_usage_metrics("") is None
|
||||
|
||||
def test_malformed(self) -> None:
|
||||
assert parse_tinfoil_usage_metrics("garbage") is None
|
||||
|
||||
def test_missing_completion(self) -> None:
|
||||
assert parse_tinfoil_usage_metrics("prompt=10") is None
|
||||
|
||||
def test_extra_whitespace(self) -> None:
|
||||
result = parse_tinfoil_usage_metrics(
|
||||
"prompt = 100 , completion = 200 , total = 300"
|
||||
)
|
||||
assert result == {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 200,
|
||||
"total_tokens": 300,
|
||||
}
|
||||
|
||||
def test_with_model_field(self) -> None:
|
||||
result = parse_tinfoil_usage_metrics(
|
||||
"prompt=42,completion=10,total=52,model=llama3-3-70b"
|
||||
)
|
||||
assert result == {
|
||||
"prompt_tokens": 42,
|
||||
"completion_tokens": 10,
|
||||
"total_tokens": 52,
|
||||
"model": "llama3-3-70b",
|
||||
}
|
||||
|
||||
def test_with_model_no_total(self) -> None:
|
||||
result = parse_tinfoil_usage_metrics(
|
||||
"prompt=67,completion=42,model=gpt-oss-120b"
|
||||
)
|
||||
assert result == {
|
||||
"prompt_tokens": 67,
|
||||
"completion_tokens": 42,
|
||||
"model": "gpt-oss-120b",
|
||||
}
|
||||
|
||||
def test_model_with_dashes_and_numbers(self) -> None:
|
||||
result = parse_tinfoil_usage_metrics(
|
||||
"prompt=1,completion=1,total=2,model=kimi-k2-6"
|
||||
)
|
||||
assert result is not None
|
||||
assert result["model"] == "kimi-k2-6"
|
||||
|
||||
def test_model_with_extra_fields(self) -> None:
|
||||
result = parse_tinfoil_usage_metrics(
|
||||
"prompt=69,completion=20,total=89,"
|
||||
"cached_prompt_tokens=64,uncached_prompt_tokens=5,"
|
||||
"model=kimi-k2-6"
|
||||
)
|
||||
assert result is not None
|
||||
assert result["prompt_tokens"] == 69
|
||||
assert result["completion_tokens"] == 20
|
||||
assert result["total_tokens"] == 89
|
||||
assert result["model"] == "kimi-k2-6"
|
||||
|
||||
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"
|
||||
)
|
||||
assert result == {
|
||||
"prompt_tokens": 67,
|
||||
"completion_tokens": 42,
|
||||
"total_tokens": 109,
|
||||
}
|
||||
assert "model" not in result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _strip_proxy_headers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStripProxyHeaders:
|
||||
def test_strips_all_proxy_only(self) -> None:
|
||||
headers = {
|
||||
"x-routstr-model": "tinfoil-llama3-3-70b",
|
||||
"X-Tinfoil-Enclave-Url": "https://inference.tinfoil.sh",
|
||||
"X-Tinfoil-Request-Usage-Metrics": "true",
|
||||
"Authorization": "Bearer secret",
|
||||
"Ehbp-Encapsulated-Key": "abc123",
|
||||
}
|
||||
clean = _strip_proxy_headers(headers)
|
||||
assert "x-routstr-model" not in clean
|
||||
assert "X-Tinfoil-Enclave-Url" not in clean
|
||||
assert "X-Tinfoil-Request-Usage-Metrics" not in clean
|
||||
assert clean["Authorization"] == "Bearer secret"
|
||||
assert clean["Ehbp-Encapsulated-Key"] == "abc123"
|
||||
|
||||
def test_all_proxy_only_headers_covered(self) -> None:
|
||||
assert _PROXY_ONLY_HEADERS == {
|
||||
"x-routstr-model",
|
||||
"x-tinfoil-enclave-url",
|
||||
"x-tinfoil-request-usage-metrics",
|
||||
}
|
||||
|
||||
|
||||
class TestPrepareEHBPUpstreamHeaders:
|
||||
def test_strips_client_proxy_headers_before_merging_target_headers(self) -> None:
|
||||
headers = {
|
||||
"x-routstr-model": "tinfoil-llama3-3-70b",
|
||||
"X-Tinfoil-Enclave-Url": "https://enclave.tinfoil.sh",
|
||||
"X-Tinfoil-Request-Usage-Metrics": "false",
|
||||
"Authorization": "Bearer upstream-key",
|
||||
"Ehbp-Encapsulated-Key": "abc123",
|
||||
}
|
||||
target_headers = {"X-Tinfoil-Request-Usage-Metrics": "true"}
|
||||
|
||||
clean = _prepare_ehbp_upstream_headers(headers, target_headers)
|
||||
|
||||
assert "x-routstr-model" not in clean
|
||||
assert "X-Tinfoil-Enclave-Url" not in clean
|
||||
assert clean["Authorization"] == "Bearer upstream-key"
|
||||
assert clean["Ehbp-Encapsulated-Key"] == "abc123"
|
||||
assert clean["X-Tinfoil-Request-Usage-Metrics"] == "true"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _resolve_ehbp_target_url
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestResolveEhbpTargetUrl:
|
||||
def test_override_with_enclave_url_for_tinfoil(self) -> None:
|
||||
result = _resolve_ehbp_target_url(
|
||||
"https://default.example.com/v1/chat/completions",
|
||||
"v1/chat/completions",
|
||||
{"X-Tinfoil-Enclave-Url": "https://enclave.tinfoil.sh"},
|
||||
"tinfoil",
|
||||
)
|
||||
assert result == "https://enclave.tinfoil.sh/v1/chat/completions"
|
||||
|
||||
def test_override_lowercase_header_for_tinfoil(self) -> None:
|
||||
result = _resolve_ehbp_target_url(
|
||||
"https://default.example.com/v1/chat/completions",
|
||||
"v1/chat/completions",
|
||||
{"x-tinfoil-enclave-url": "https://enclave.tinfoil.sh"},
|
||||
"tinfoil",
|
||||
)
|
||||
assert result == "https://enclave.tinfoil.sh/v1/chat/completions"
|
||||
|
||||
def test_no_override(self) -> None:
|
||||
default = "https://inference.tinfoil.sh/v1/chat/completions"
|
||||
result = _resolve_ehbp_target_url(
|
||||
default,
|
||||
"v1/chat/completions",
|
||||
{},
|
||||
"tinfoil",
|
||||
)
|
||||
assert result == default
|
||||
|
||||
def test_non_tinfoil_provider_ignores_enclave_url(self) -> None:
|
||||
default = "https://api.ppq.ai/private/v1/chat/completions"
|
||||
result = _resolve_ehbp_target_url(
|
||||
default,
|
||||
"v1/chat/completions",
|
||||
{"X-Tinfoil-Enclave-Url": "https://enclave.tinfoil.sh"},
|
||||
"ppqai",
|
||||
)
|
||||
assert result == default
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_url",
|
||||
[
|
||||
"http://enclave.tinfoil.sh",
|
||||
"https://attacker.example",
|
||||
"https://tinfoil.sh.attacker.example",
|
||||
"https://127.0.0.1",
|
||||
"https://enclave.tinfoil.sh:8443",
|
||||
"https://user:pass@enclave.tinfoil.sh",
|
||||
],
|
||||
)
|
||||
def test_tinfoil_rejects_unsafe_enclave_url(self, bad_url: str) -> None:
|
||||
from routstr.core.exceptions import UpstreamError
|
||||
|
||||
with pytest.raises(UpstreamError):
|
||||
_resolve_ehbp_target_url(
|
||||
"https://default.example.com/v1/chat/completions",
|
||||
"v1/chat/completions",
|
||||
{"X-Tinfoil-Enclave-Url": bad_url},
|
||||
"tinfoil",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _compute_ehbp_actual_cost
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestComputeEhbpActualCost:
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_usage_falls_back_to_max_cost(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["input_tokens"] == 0
|
||||
assert result["output_tokens"] == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usage_parsed_and_clamped(self) -> None:
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "llama3-3-70b"
|
||||
model_obj.forwarded_model_id = "llama3-3-70b"
|
||||
# The actual cost from calculate_cost will be small; we just verify
|
||||
# it's clamped to min_request_msat at minimum.
|
||||
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=10,
|
||||
output_msats=20,
|
||||
total_msats=30,
|
||||
total_usd=0.0001,
|
||||
input_tokens=67,
|
||||
output_tokens=42,
|
||||
)
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=67,completion=42,total=109",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
assert result["total_msats"] == 30
|
||||
assert result["total_msats"] <= 100_000
|
||||
assert result["input_tokens"] == 67
|
||||
assert result["output_tokens"] == 42
|
||||
assert result["total_tokens"] == 109
|
||||
assert result["input_msats"] == 10
|
||||
assert result["output_msats"] == 20
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_cost_data_falls_back(self) -> None:
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "llama3-3-70b"
|
||||
model_obj.forwarded_model_id = "llama3-3-70b"
|
||||
with patch(
|
||||
"routstr.upstream.ehbp.calculate_cost",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_calc:
|
||||
from routstr.payment.cost_calculation import MaxCostData
|
||||
|
||||
mock_calc.return_value = MaxCostData(
|
||||
base_msats=0,
|
||||
input_msats=0,
|
||||
output_msats=0,
|
||||
total_msats=0,
|
||||
total_usd=0.0,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
)
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=0,completion=0",
|
||||
model_obj,
|
||||
50_000,
|
||||
)
|
||||
assert result["total_msats"] == 50_000
|
||||
assert result["input_tokens"] == 0
|
||||
assert result["output_tokens"] == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_match_no_actual_model_key(self) -> None:
|
||||
"""When the served model matches the requested one, no actual_model key."""
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "llama3-3-70b"
|
||||
model_obj.forwarded_model_id = "llama3-3-70b"
|
||||
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=10,
|
||||
total_msats=15,
|
||||
total_usd=0.0,
|
||||
input_tokens=42,
|
||||
output_tokens=10,
|
||||
)
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=42,completion=10,total=52,model=llama3-3-70b",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
assert "actual_model" not in result
|
||||
# calculate_cost called with requested model
|
||||
call_args = mock_calc.call_args
|
||||
assert call_args[0][0]["model"] == "llama3-3-70b"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_alias_match_no_actual_model_key(self) -> None:
|
||||
"""When the served upstream model matches forwarded_model_id through
|
||||
a client-facing alias, no actual_model key is set."""
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "tinfoil-glm-5-2" # client-facing alias
|
||||
model_obj.forwarded_model_id = "glm-5-2" # actual upstream ID
|
||||
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=10,
|
||||
total_msats=15,
|
||||
total_usd=0.0,
|
||||
input_tokens=42,
|
||||
output_tokens=10,
|
||||
)
|
||||
# Tinfoil header returns the actual upstream model ID
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=42,completion=10,total=52,model=glm-5-2",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
assert "actual_model" not in result
|
||||
# calculate_cost called with the client-facing model ID (whose
|
||||
# pricing includes the correct upstream rates)
|
||||
call_args = mock_calc.call_args
|
||||
assert call_args[0][0]["model"] == "tinfoil-glm-5-2"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_real_mismatch_uses_actual_model_for_pricing(self) -> None:
|
||||
"""When the served model differs from the expected upstream model,
|
||||
the actual model's pricing is used."""
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "tinfoil-gpt-oss-120b" # client-facing alias
|
||||
model_obj.forwarded_model_id = "gpt-oss-120b" # expected upstream
|
||||
|
||||
actual_model_obj = MagicMock()
|
||||
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:
|
||||
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,
|
||||
)
|
||||
# Tinfoil served llama3-3-70b instead of gpt-oss-120b
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=42,completion=10,total=52,model=llama3-3-70b",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
assert result["actual_model"] == "llama3-3-70b"
|
||||
assert result["total_msats"] == 60
|
||||
# calculate_cost called with the actual model's client-facing ID
|
||||
call_args = mock_calc.call_args
|
||||
assert call_args[0][0]["model"] == "tinfoil-llama3-3-70b"
|
||||
|
||||
@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."""
|
||||
model_obj = MagicMock()
|
||||
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:
|
||||
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=42,
|
||||
output_tokens=10,
|
||||
)
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=42,completion=10,total=52,model=nonexistent",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
assert "actual_model" not in result
|
||||
# calculate_cost called with the requested model (fallback)
|
||||
call_args = mock_calc.call_args
|
||||
assert call_args[0][0]["model"] == "gpt-oss-120b"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_old_format_no_model_uses_requested(self) -> None:
|
||||
"""Old format without model field uses requested model for pricing."""
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "llama3-3-70b"
|
||||
model_obj.forwarded_model_id = "llama3-3-70b"
|
||||
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=10,
|
||||
total_msats=15,
|
||||
total_usd=0.0,
|
||||
input_tokens=67,
|
||||
output_tokens=42,
|
||||
)
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=67,completion=42,total=109",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
assert "actual_model" not in result
|
||||
call_args = mock_calc.call_args
|
||||
assert call_args[0][0]["model"] == "llama3-3-70b"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_case_insensitive_model_match(self) -> None:
|
||||
"""Casing differences between the header and forwarded_model_id
|
||||
should not trigger a spurious mismatch."""
|
||||
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:
|
||||
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=42,
|
||||
output_tokens=10,
|
||||
)
|
||||
# Header returns uppercase — same model, different casing
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=42,completion=10,total=52,model=GLM-5-2",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
assert "actual_model" not in result
|
||||
# No mismatch: requested model pricing used
|
||||
call_args = mock_calc.call_args
|
||||
assert call_args[0][0]["model"] == "tinfoil-glm-5-2"
|
||||
mock_get_model.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_date_versioned_alias_resolves_to_requested(self) -> None:
|
||||
"""When the served model is a date-versioned alias that resolves back
|
||||
to the requested model, no mismatch is propagated."""
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "tinfoil-glm-5-2"
|
||||
model_obj.forwarded_model_id = "glm-5-2"
|
||||
|
||||
resolved_model_obj = MagicMock()
|
||||
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:
|
||||
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=42,
|
||||
output_tokens=10,
|
||||
)
|
||||
# Tinfoil returns a date-versioned ID with different casing.
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=42,completion=10,total=52,model=GLM-5-2-20260415",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
# Registry resolution, rather than unconditional suffix removal,
|
||||
# establishes that this alias represents the expected model.
|
||||
assert "actual_model" not in result
|
||||
assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2"
|
||||
mock_get_model.assert_called_once_with("GLM-5-2-20260415")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_configured_date_version_is_preserved_as_identity(self) -> None:
|
||||
"""A date suffix in forwarded_model_id is meaningful and preserved."""
|
||||
model_obj = MagicMock()
|
||||
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:
|
||||
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=42,
|
||||
output_tokens=10,
|
||||
)
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=42,completion=10,total=52,model=GLM-5-2-20260415",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
|
||||
assert "actual_model" not in result
|
||||
assert (
|
||||
mock_calc.call_args[0][0]["model"]
|
||||
== "tinfoil-glm-5-2-20260415"
|
||||
)
|
||||
mock_get_model.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_different_client_alias_same_upstream_identity(self) -> None:
|
||||
"""A global alias winner from another provider is not a failover when
|
||||
its forwarded model ID matches the requested upstream identity."""
|
||||
model_obj = MagicMock()
|
||||
model_obj.id = "tinfoil-glm-5-2"
|
||||
model_obj.forwarded_model_id = "glm-5-2"
|
||||
|
||||
resolved_model_obj = MagicMock()
|
||||
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:
|
||||
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=42,
|
||||
output_tokens=10,
|
||||
)
|
||||
result = await _compute_ehbp_actual_cost(
|
||||
"prompt=42,completion=10,total=52,model=provider-alias",
|
||||
model_obj,
|
||||
100_000,
|
||||
)
|
||||
|
||||
mock_get_model.assert_called_once_with("provider-alias")
|
||||
assert "actual_model" not in result
|
||||
assert mock_calc.call_args[0][0]["model"] == "tinfoil-glm-5-2"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TinfoilUpstreamProvider
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTinfoilUpstreamProvider:
|
||||
def test_provider_type_and_defaults(self) -> None:
|
||||
assert TinfoilUpstreamProvider.provider_type == "tinfoil"
|
||||
assert (
|
||||
TinfoilUpstreamProvider.default_base_url
|
||||
== "https://inference.tinfoil.sh"
|
||||
)
|
||||
assert TinfoilUpstreamProvider.supports_ehbp is True
|
||||
|
||||
def test_transform_model_name(self) -> None:
|
||||
provider = TinfoilUpstreamProvider(api_key="test")
|
||||
assert provider.transform_model_name("tinfoil/llama3-3-70b") == "llama3-3-70b"
|
||||
assert provider.transform_model_name("llama3-3-70b") == "llama3-3-70b"
|
||||
|
||||
def test_get_ehbp_forwarding_target_includes_usage_header(self) -> None:
|
||||
provider = TinfoilUpstreamProvider(api_key="test")
|
||||
model_obj = MagicMock()
|
||||
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 "v1/chat/completions" in target.url
|
||||
|
||||
def test_get_provider_metadata(self) -> None:
|
||||
meta = TinfoilUpstreamProvider.get_provider_metadata()
|
||||
assert meta["id"] == "tinfoil"
|
||||
assert meta["name"] == "Tinfoil"
|
||||
assert meta["fixed_base_url"] is True
|
||||
|
||||
def test_tinfoil_model_pricing_parses(self) -> None:
|
||||
data = {
|
||||
"id": "llama3-3-70b",
|
||||
"context_window": 128000,
|
||||
"created": 1721764788,
|
||||
"pricing": {
|
||||
"inputTokenPricePer1M": 1.75,
|
||||
"outputTokenPricePer1M": 2.75,
|
||||
"requestPrice": 0,
|
||||
},
|
||||
"endpoints": ["/v1/chat/completions"],
|
||||
"type": "chat",
|
||||
}
|
||||
tf = TinfoilModel.parse_obj(data)
|
||||
assert tf.id == "llama3-3-70b"
|
||||
assert tf.pricing.inputTokenPricePer1M == 1.75
|
||||
assert tf.pricing.outputTokenPricePer1M == 2.75
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_models_parses_response(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": "llama3-3-70b",
|
||||
"context_window": 128000,
|
||||
"created": 1721764788,
|
||||
"multimodal": False,
|
||||
"pricing": {
|
||||
"inputTokenPricePer1M": 1.75,
|
||||
"outputTokenPricePer1M": 2.75,
|
||||
"requestPrice": 0,
|
||||
},
|
||||
"endpoints": ["/v1/chat/completions"],
|
||||
"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].id == "llama3-3-70b"
|
||||
assert models[0].pricing.prompt == 1.75 / 1_000_000
|
||||
assert models[0].pricing.completion == 2.75 / 1_000_000
|
||||
assert models[0].context_length == 128000
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_models_handles_error(self) -> None:
|
||||
provider = TinfoilUpstreamProvider(api_key="test")
|
||||
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(side_effect=Exception("network error"))
|
||||
mock_client_cls.return_value = mock_client
|
||||
|
||||
models = await provider.fetch_models()
|
||||
|
||||
assert models == []
|
||||
@@ -0,0 +1,138 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.upstream.tinfoil_trailer import forward_with_trailer
|
||||
|
||||
|
||||
class FakeReader:
|
||||
def __init__(self, chunks: list[bytes]) -> None:
|
||||
self._chunks = chunks
|
||||
|
||||
async def read(self, _size: int) -> bytes:
|
||||
if self._chunks:
|
||||
return self._chunks.pop(0)
|
||||
return b""
|
||||
|
||||
|
||||
class FakeWriter:
|
||||
def __init__(self) -> None:
|
||||
self.written = b""
|
||||
self.drain = AsyncMock()
|
||||
self.wait_closed = AsyncMock()
|
||||
self.close = MagicMock()
|
||||
|
||||
def write(self, data: bytes) -> None:
|
||||
self.written += data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forward_with_trailer_captures_usage_trailer(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
response = (
|
||||
b"HTTP/1.1 200 OK\r\n"
|
||||
b"Transfer-Encoding: chunked\r\n"
|
||||
b"Trailer: X-Tinfoil-Usage-Metrics\r\n"
|
||||
b"\r\n"
|
||||
b"5\r\nhello\r\n"
|
||||
b"0\r\n"
|
||||
b"X-Tinfoil-Usage-Metrics: prompt=1,completion=2,total=3\r\n"
|
||||
b"\r\n"
|
||||
)
|
||||
reader = FakeReader([response])
|
||||
writer = FakeWriter()
|
||||
open_connection = AsyncMock(return_value=(reader, writer))
|
||||
monkeypatch.setattr(
|
||||
"routstr.upstream.tinfoil_trailer.asyncio.open_connection", open_connection
|
||||
)
|
||||
|
||||
result = await forward_with_trailer(
|
||||
method="POST",
|
||||
url="https://enclave.tinfoil.sh/v1/chat/completions?stream=true",
|
||||
headers={"Authorization": "Bearer upstream"},
|
||||
body=b"opaque",
|
||||
)
|
||||
|
||||
assert result.status_code == 200
|
||||
assert result.body == b"hello"
|
||||
assert result.trailers == [
|
||||
("x-tinfoil-usage-metrics", "prompt=1,completion=2,total=3")
|
||||
]
|
||||
assert b"Connection: close" in writer.written
|
||||
writer.close.assert_called_once()
|
||||
writer.wait_closed.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forward_with_trailer_strips_hop_by_hop_headers(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
response = b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok"
|
||||
reader = FakeReader([response])
|
||||
writer = FakeWriter()
|
||||
monkeypatch.setattr(
|
||||
"routstr.upstream.tinfoil_trailer.asyncio.open_connection",
|
||||
AsyncMock(return_value=(reader, writer)),
|
||||
)
|
||||
|
||||
await forward_with_trailer(
|
||||
method="POST",
|
||||
url="https://enclave.tinfoil.sh/v1/chat/completions",
|
||||
headers={
|
||||
"Authorization": "Bearer upstream",
|
||||
"Connection": "keep-alive, X-Client-Hop",
|
||||
"Keep-Alive": "timeout=5",
|
||||
"Proxy-Authenticate": "Basic",
|
||||
"Proxy-Authorization": "Basic secret",
|
||||
"TE": "trailers",
|
||||
"Trailer": "X-Usage",
|
||||
"Transfer-Encoding": "chunked",
|
||||
"Upgrade": "websocket",
|
||||
"X-Client-Hop": "remove-me",
|
||||
"X-End-To-End": "preserve-me",
|
||||
},
|
||||
body=b"opaque",
|
||||
)
|
||||
|
||||
serialized_headers = writer.written.split(b"\r\n\r\n", 1)[0].lower()
|
||||
for name in (
|
||||
b"keep-alive",
|
||||
b"proxy-authenticate",
|
||||
b"proxy-authorization",
|
||||
b"te:",
|
||||
b"trailer:",
|
||||
b"transfer-encoding",
|
||||
b"upgrade:",
|
||||
b"x-client-hop",
|
||||
):
|
||||
assert name not in serialized_headers
|
||||
assert b"connection: close" in serialized_headers
|
||||
assert b"content-length: 6" in serialized_headers
|
||||
assert b"x-end-to-end: preserve-me" in serialized_headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forward_with_trailer_enforces_response_size_limit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
response = b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\n\r\nhello"
|
||||
reader = FakeReader([response])
|
||||
writer = FakeWriter()
|
||||
monkeypatch.setattr(
|
||||
"routstr.upstream.tinfoil_trailer.asyncio.open_connection",
|
||||
AsyncMock(return_value=(reader, writer)),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="EHBP response exceeded"):
|
||||
await forward_with_trailer(
|
||||
method="POST",
|
||||
url="https://enclave.tinfoil.sh/v1/chat/completions",
|
||||
headers={},
|
||||
body=b"opaque",
|
||||
max_response_bytes=4,
|
||||
)
|
||||
|
||||
writer.close.assert_called_once()
|
||||
@@ -362,8 +362,11 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None:
|
||||
revert_mock = AsyncMock(return_value=True)
|
||||
|
||||
with (
|
||||
patch.object(proxy_module, "get_model_instance", return_value=MagicMock()),
|
||||
patch.object(proxy_module, "get_provider_for_model", return_value=[upstream]),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"get_candidates",
|
||||
return_value=[(MagicMock(), upstream)],
|
||||
),
|
||||
patch.object(
|
||||
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
|
||||
),
|
||||
|
||||
@@ -0,0 +1,339 @@
|
||||
"""Tests for ``routstr.core.vault`` — the secret encrypt/hash/fingerprint helpers.
|
||||
|
||||
Specifies the primitives that the rest of the secret-storage work (issue #553)
|
||||
builds on, independent of any database or app wiring:
|
||||
|
||||
- ``encrypt``/``decrypt`` — Fernet symmetric encryption emitting self-describing
|
||||
``fernet:v1:`` ciphertext, so a value can be told apart from legacy plaintext
|
||||
and from ciphertext written under a different ``ROUTSTR_SECRET_KEY`` (which
|
||||
surfaces as a hard ``InvalidToken`` rather than silent corruption).
|
||||
- ``hash_password``/``verify_password`` — salted scrypt hashing that is
|
||||
*key-independent* (does not depend on ``ROUTSTR_SECRET_KEY``), so password
|
||||
login and the recovery script keep working even if the key is lost.
|
||||
- a missing/malformed ``ROUTSTR_SECRET_KEY`` fails fast with the generation
|
||||
command in the message.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from cryptography.fernet import InvalidToken
|
||||
|
||||
from routstr.core import vault
|
||||
|
||||
# Two distinct, valid Fernet keys held fixed so ciphertext/fingerprints are
|
||||
# reproducible across runs and we can exercise the wrong-key path.
|
||||
KEY_A = "l_Tkp-7xmjcQ-IFhr6qhILrU8HPRbEmYMrfSbo_5srU="
|
||||
KEY_B = "_Teyrky_iToeDK51Tj1FsI9MJ340_cqKGmeher-a7MQ="
|
||||
|
||||
|
||||
def _use_key(monkeypatch: pytest.MonkeyPatch, key: str) -> None:
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY", key)
|
||||
|
||||
|
||||
# --- encrypt / decrypt -----------------------------------------------------
|
||||
|
||||
|
||||
def test_encrypt_decrypt_round_trips(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
assert vault.decrypt(vault.encrypt("nsec1secret")) == "nsec1secret"
|
||||
|
||||
|
||||
def test_encrypt_emits_self_describing_prefix(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
assert vault.encrypt("x").startswith("fernet:v1:")
|
||||
|
||||
|
||||
def test_encrypt_is_non_deterministic(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
# Fernet embeds a random IV/timestamp: equal plaintext -> different
|
||||
# ciphertext. This is exactly why upstream-key equality needs a blind index.
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
assert vault.encrypt("same") != vault.encrypt("same")
|
||||
|
||||
|
||||
def test_is_encrypted_distinguishes_ciphertext_from_plaintext(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
assert vault.is_encrypted(vault.encrypt("x")) is True
|
||||
assert vault.is_encrypted("sk-plaintext-api-key") is False
|
||||
assert vault.is_encrypted("") is False
|
||||
|
||||
|
||||
def test_decrypt_rejects_unprefixed_value(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# Guards the migration paths: a legacy plaintext value must never be
|
||||
# mistaken for ciphertext and "decrypted".
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
with pytest.raises(ValueError):
|
||||
vault.decrypt("not-encrypted")
|
||||
|
||||
|
||||
def test_decrypt_with_wrong_key_raises(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# The fail-fast signal: ciphertext written under KEY_A cannot be read under
|
||||
# KEY_B -> InvalidToken (bootstrap turns this into a clear startup error).
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
token = vault.encrypt("secret")
|
||||
_use_key(monkeypatch, KEY_B)
|
||||
with pytest.raises(InvalidToken):
|
||||
vault.decrypt(token)
|
||||
|
||||
|
||||
# --- password hashing (key-independent) ------------------------------------
|
||||
|
||||
|
||||
def test_hash_and_verify_password(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
stored = vault.hash_password("correct horse")
|
||||
assert vault.verify_password("correct horse", stored) is True
|
||||
assert vault.verify_password("wrong", stored) is False
|
||||
|
||||
|
||||
def test_password_hash_is_salted(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
a = vault.hash_password("pw")
|
||||
b = vault.hash_password("pw")
|
||||
assert a != b
|
||||
assert vault.verify_password("pw", a) is True
|
||||
assert vault.verify_password("pw", b) is True
|
||||
|
||||
|
||||
def test_verify_password_rejects_malformed_stored_value(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# A garbage or non-scrypt stored value must verify to False, never raise.
|
||||
_use_key(monkeypatch, KEY_A)
|
||||
assert vault.verify_password("pw", "") is False
|
||||
assert vault.verify_password("pw", "not-a-hash") is False
|
||||
assert vault.verify_password("pw", "bcrypt:1:2:3:x:y") is False
|
||||
|
||||
|
||||
def test_password_hashing_is_key_independent(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
# scrypt does not use ROUTSTR_SECRET_KEY, so login and the recovery script
|
||||
# work even when the key is missing.
|
||||
monkeypatch.delenv("ROUTSTR_SECRET_KEY", raising=False)
|
||||
stored = vault.hash_password("pw")
|
||||
assert vault.verify_password("pw", stored) is True
|
||||
|
||||
|
||||
# --- fail-fast on missing/malformed key ------------------------------------
|
||||
|
||||
|
||||
def test_decrypt_without_any_key_fails_fast_with_generation_command(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
# Reading secrets is strict: with no key in env AND no key file, decrypt
|
||||
# fails fast with the generation command rather than silently minting a new
|
||||
# key (a fresh key could never match already-encrypted ciphertext). Only the
|
||||
# encrypt path auto-provisions; the read path never does.
|
||||
monkeypatch.delenv("ROUTSTR_SECRET_KEY", raising=False)
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY_FILE", str(tmp_path / "absent.key"))
|
||||
with pytest.raises(RuntimeError) as exc:
|
||||
vault.decrypt("fernet:v1:not-real-ciphertext")
|
||||
msg = str(exc.value)
|
||||
assert "ROUTSTR_SECRET_KEY" in msg
|
||||
assert "Fernet.generate_key" in msg
|
||||
|
||||
|
||||
def test_malformed_key_fails_fast(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY", "not-a-valid-fernet-key")
|
||||
with pytest.raises(RuntimeError):
|
||||
vault.encrypt("x")
|
||||
|
||||
|
||||
def test_malformed_env_key_does_not_self_provision(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
# A malformed env key is an operator mistake, not an unset key: it must fail
|
||||
# loudly, never silently generate a different key to a file (which would hide
|
||||
# the mistake and could brick secrets the operator meant to key differently).
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY", "not-a-valid-fernet-key")
|
||||
key_file = tmp_path / "routstr_secret.key"
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY_FILE", str(key_file))
|
||||
with pytest.raises(RuntimeError):
|
||||
vault.encrypt("x")
|
||||
assert not key_file.exists()
|
||||
|
||||
|
||||
# --- auto-provisioned key file (non-breaking upgrade path) -----------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def generated_key_file(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> Path:
|
||||
"""No env key; the key file points at a fresh, empty tmp location.
|
||||
|
||||
Exercises what an existing node hits when it upgrades without setting
|
||||
ROUTSTR_SECRET_KEY: the master key is auto-generated and persisted here so
|
||||
boot does not break, while secrets are still never written in plaintext.
|
||||
"""
|
||||
monkeypatch.delenv("ROUTSTR_SECRET_KEY", raising=False)
|
||||
key_file = tmp_path / "routstr_secret.key"
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY_FILE", str(key_file))
|
||||
return key_file
|
||||
|
||||
|
||||
def test_encrypt_without_key_generates_and_persists_key_file(
|
||||
generated_key_file: Path,
|
||||
) -> None:
|
||||
# Encryption at rest stays mandatory, but a missing key is provisioned rather
|
||||
# than fatal: encrypt generates a key, writes it to the key file (owner-only),
|
||||
# and the value round-trips — decrypt, still with no env key, reads the same
|
||||
# file key back.
|
||||
key_file = generated_key_file
|
||||
assert not key_file.exists()
|
||||
|
||||
token = vault.encrypt("nsec1secret")
|
||||
|
||||
assert token.startswith("fernet:v1:")
|
||||
assert key_file.exists()
|
||||
assert key_file.stat().st_mode & 0o077 == 0 # not group/other-accessible
|
||||
assert vault.decrypt(token) == "nsec1secret"
|
||||
|
||||
|
||||
def test_generated_key_warns_operator_with_path_not_value(
|
||||
generated_key_file: Path, capsys: pytest.CaptureFixture[str]
|
||||
) -> None:
|
||||
# The notice names the file to back up and shouts the back-up imperative so an
|
||||
# upgrading operator cannot miss it, but it MUST NOT echo the key value: the
|
||||
# secret lives in the 0600 file, and printing it would leak it into captured
|
||||
# stdout / aggregated container logs.
|
||||
key_file = generated_key_file
|
||||
vault.encrypt("x")
|
||||
|
||||
out = capsys.readouterr().out
|
||||
assert "ROUTSTR_SECRET_KEY" in out
|
||||
assert str(key_file) in out
|
||||
assert key_file.read_text().strip() not in out
|
||||
assert "BACK UP" in out.upper()
|
||||
|
||||
|
||||
def test_existing_key_file_is_reused_and_warns_only_once(
|
||||
generated_key_file: Path, capsys: pytest.CaptureFixture[str]
|
||||
) -> None:
|
||||
# Once the key exists, later encrypts reuse it (never rotate a key that
|
||||
# secrets were already encrypted under) and stay silent (no repeated notice).
|
||||
key_file = generated_key_file
|
||||
vault.encrypt("a")
|
||||
first_key = key_file.read_text()
|
||||
capsys.readouterr() # drain the one-time notice
|
||||
|
||||
vault.encrypt("b")
|
||||
|
||||
assert key_file.read_text() == first_key
|
||||
assert capsys.readouterr().out == ""
|
||||
|
||||
|
||||
def test_generated_key_is_published_atomically(
|
||||
generated_key_file: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
# The key must appear at its final path only as a complete file: it is written
|
||||
# to a temp file and atomically linked into place. If the publish (link) step
|
||||
# fails — e.g. the process crashes — the destination must be absent, never a
|
||||
# half-written or empty file that a later boot would read as a corrupt key and
|
||||
# then refuse to decrypt every secret. No temp debris is left behind.
|
||||
key_file = generated_key_file
|
||||
|
||||
def boom(src: object, dst: object) -> None:
|
||||
raise OSError("crash during atomic publish")
|
||||
|
||||
monkeypatch.setattr(vault.os, "link", boom)
|
||||
|
||||
with pytest.raises(OSError):
|
||||
vault.encrypt("x")
|
||||
|
||||
assert not key_file.exists()
|
||||
assert list(key_file.parent.iterdir()) == []
|
||||
|
||||
|
||||
def test_racing_worker_adopts_winners_key_without_clobber(
|
||||
generated_key_file: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
# Two workers auto-generate a key concurrently. The first to link "wins" and
|
||||
# its key is the one on disk; a worker that loses the link race must adopt the
|
||||
# winner's key — secrets may already be encrypted under it — rather than
|
||||
# clobber it or crash. Force the race deterministically: the winner publishes
|
||||
# its key at the final path just as this worker tries to link, so os.link
|
||||
# raises FileExistsError and the loser reads the winner's key back.
|
||||
key_file = generated_key_file
|
||||
|
||||
def winner_links_first(src: object, dst: object) -> None:
|
||||
key_file.write_text(KEY_A) # the winner's already-published key
|
||||
raise FileExistsError
|
||||
|
||||
monkeypatch.setattr(vault.os, "link", winner_links_first)
|
||||
|
||||
token = vault.encrypt("secret")
|
||||
|
||||
# The winner's key stays put and the loser encrypted under it, so the value
|
||||
# round-trips under KEY_A even though this worker had generated its own key.
|
||||
assert key_file.read_text().strip() == KEY_A
|
||||
assert vault.decrypt(token) == "secret"
|
||||
# The losing worker left no temp debris behind.
|
||||
assert [p.name for p in key_file.parent.iterdir()] == [key_file.name]
|
||||
|
||||
|
||||
def test_loose_key_file_perms_are_tightened_on_read(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
# A key file left group/other-readable (e.g. written under a loose umask
|
||||
# before this hardening, or by a careless operator) is a leaked-secret risk.
|
||||
# Reading it repairs the permissions to owner-only rather than trusting a
|
||||
# world-readable master key, while still using the key so boot is not broken.
|
||||
monkeypatch.delenv("ROUTSTR_SECRET_KEY", raising=False)
|
||||
key_file = tmp_path / "routstr_secret.key"
|
||||
key_file.write_text(KEY_A)
|
||||
key_file.chmod(0o644)
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY_FILE", str(key_file))
|
||||
|
||||
token = vault.encrypt("secret") # reads the loose file, repairs its perms
|
||||
|
||||
assert key_file.stat().st_mode & 0o077 == 0
|
||||
assert vault.decrypt(token) == "secret"
|
||||
|
||||
|
||||
def test_env_key_takes_precedence_over_key_file(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
# An explicit env key wins over a persisted file key (an operator-supplied
|
||||
# key from a secrets manager overrides the auto-generated one), and the file
|
||||
# is left untouched.
|
||||
key_file = tmp_path / "routstr_secret.key"
|
||||
key_file.write_text(KEY_B)
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY_FILE", str(key_file))
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY", KEY_A)
|
||||
|
||||
token = vault.encrypt("x")
|
||||
|
||||
assert vault.decrypt(token) == "x" # env key (A) decrypts it
|
||||
monkeypatch.setenv("ROUTSTR_SECRET_KEY", KEY_B)
|
||||
with pytest.raises(InvalidToken):
|
||||
vault.decrypt(token) # the file key (B) does not
|
||||
assert key_file.read_text() == KEY_B # file key never used or overwritten
|
||||
|
||||
|
||||
def test_generated_key_defaults_beside_the_database(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
|
||||
) -> None:
|
||||
# With no env key and no explicit key-file path, the key is provisioned next
|
||||
# to the SQLite database, so it rides whatever volume already persists the
|
||||
# data instead of landing in the working directory (where a container
|
||||
# recreate would lose it and brick decryption).
|
||||
monkeypatch.delenv("ROUTSTR_SECRET_KEY", raising=False)
|
||||
monkeypatch.delenv("ROUTSTR_SECRET_KEY_FILE", raising=False)
|
||||
monkeypatch.chdir(tmp_path) # isolate the working-dir fallback from the repo
|
||||
db_dir = tmp_path / "data"
|
||||
db_dir.mkdir()
|
||||
monkeypatch.setenv("DATABASE_URL", f"sqlite+aiosqlite:///{db_dir}/routstr.db")
|
||||
|
||||
token = vault.encrypt("beside-the-db")
|
||||
|
||||
assert (db_dir / "routstr_secret.key").exists()
|
||||
assert not (tmp_path / "routstr_secret.key").exists() # not the working dir
|
||||
assert vault.decrypt(token) == "beside-the-db"
|
||||
@@ -15,6 +15,7 @@ from routstr.wallet import (
|
||||
get_balance,
|
||||
is_mint_connection_error,
|
||||
recieve_token,
|
||||
send,
|
||||
send_token,
|
||||
)
|
||||
|
||||
@@ -26,7 +27,11 @@ async def test_get_balance() -> None:
|
||||
mock_wallet.load_mint = AsyncMock()
|
||||
mock_wallet.load_proofs = AsyncMock()
|
||||
|
||||
with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet):
|
||||
# Reset the module-level wallet cache so a real wallet cached by an earlier
|
||||
# test (e.g. an unmocked admin-withdraw path) can't shadow the mock here.
|
||||
with patch("routstr.wallet._wallets", {}), patch(
|
||||
"routstr.wallet.Wallet.with_db", return_value=mock_wallet
|
||||
):
|
||||
balance = await get_balance("sat")
|
||||
assert balance == 50000
|
||||
|
||||
@@ -147,6 +152,69 @@ async def test_send_token() -> None:
|
||||
assert token == "test_token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_falls_back_when_preferred_mint_has_only_reserved_balance() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
preferred_wallet = Mock(keysets={}, proofs=[])
|
||||
preferred_wallet.select_to_send = AsyncMock()
|
||||
primary_wallet = Mock(keysets={}, proofs=[])
|
||||
primary_wallet.select_to_send = AsyncMock()
|
||||
primary_wallet.serialize_proofs = AsyncMock(return_value="primary-token")
|
||||
primary_wallet.set_reserved_for_send = AsyncMock()
|
||||
|
||||
preferred_liquid = Mock(amount=500, reserved=False)
|
||||
preferred_reserved = Mock(amount=600, reserved=True)
|
||||
primary_liquid = Mock(amount=1000, reserved=False)
|
||||
primary_wallet.select_to_send.return_value = ([primary_liquid], None)
|
||||
|
||||
async def get_wallet(mint_url: str, unit: str) -> Mock:
|
||||
assert unit == "sat"
|
||||
return primary_wallet if mint_url == "http://primary:3338" else preferred_wallet
|
||||
|
||||
def get_proofs(wallet: Mock, mint_url: str, unit: str) -> list[Mock]:
|
||||
assert unit == "sat"
|
||||
if wallet is primary_wallet:
|
||||
assert mint_url == "http://primary:3338"
|
||||
return [primary_liquid]
|
||||
assert mint_url == "http://preferred:3338"
|
||||
return [preferred_liquid, preferred_reserved]
|
||||
|
||||
with (
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
patch("routstr.wallet.get_wallet", side_effect=get_wallet),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", side_effect=get_proofs),
|
||||
):
|
||||
amount, token = await send(1000, "sat", "http://preferred:3338")
|
||||
|
||||
assert (amount, token) == (1000, "primary-token")
|
||||
preferred_wallet.select_to_send.assert_not_awaited()
|
||||
primary_wallet.select_to_send.assert_awaited_once_with(
|
||||
[primary_liquid], 1000, set_reserved=False, include_fees=False
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_primary_with_only_reserved_proofs_still_raises() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
wallet = Mock(keysets={}, proofs=[])
|
||||
wallet.select_to_send = AsyncMock(side_effect=RuntimeError("balance too low"))
|
||||
reserved = Mock(amount=1000, reserved=True)
|
||||
|
||||
with (
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[reserved]),
|
||||
pytest.raises(RuntimeError, match="balance too low"),
|
||||
):
|
||||
await send(1000, "sat", "http://primary:3338")
|
||||
|
||||
wallet.select_to_send.assert_awaited_once_with(
|
||||
[], 1000, set_reserved=False, include_fees=False
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_credit_balance() -> None:
|
||||
token_data = {
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
"""Additional money-path coverage tests for wallet.py (86% → target 90%).
|
||||
|
||||
Tests error classification, periodic task structure, and token operations.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
# ===========================================================================
|
||||
# is_mint_connection_error
|
||||
# ===========================================================================
|
||||
|
||||
def test_is_mint_connection_error_true() -> None:
|
||||
"""Connection errors are detected."""
|
||||
from routstr.wallet import is_mint_connection_error
|
||||
|
||||
assert is_mint_connection_error(ConnectionRefusedError("refused")) is True
|
||||
assert is_mint_connection_error(TimeoutError("timeout")) is True
|
||||
|
||||
|
||||
def test_is_mint_connection_error_false() -> None:
|
||||
"""Non-connection errors are not flagged."""
|
||||
from routstr.wallet import is_mint_connection_error
|
||||
|
||||
assert is_mint_connection_error(ValueError("bad data")) is False
|
||||
assert is_mint_connection_error(KeyError("missing key")) is False
|
||||
assert is_mint_connection_error(RuntimeError("something broke")) is False
|
||||
assert is_mint_connection_error(AttributeError("no attr")) is False
|
||||
# OSError is NOT a connection error unless it's a subclass
|
||||
assert is_mint_connection_error(OSError("generic")) is False
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# classify_redemption_error
|
||||
# ===========================================================================
|
||||
|
||||
def test_classify_redemption_error_token_consumed() -> None:
|
||||
"""Token already spent returns token_consumed classification."""
|
||||
from routstr.wallet import TokenConsumedError, classify_redemption_error
|
||||
|
||||
result = classify_redemption_error(
|
||||
TokenConsumedError("Token was already redeemed")
|
||||
)
|
||||
assert result is not None
|
||||
assert result[0] == "token_consumed"
|
||||
assert result[1] == 500
|
||||
|
||||
|
||||
def test_classify_redemption_error_mint_connection() -> None:
|
||||
"""Mint connection error is classified correctly."""
|
||||
from routstr.wallet import classify_redemption_error
|
||||
|
||||
result = classify_redemption_error(
|
||||
ConnectionRefusedError("Connection refused")
|
||||
)
|
||||
assert result is not None
|
||||
# Should classify as mint_connection or return error tuple
|
||||
assert isinstance(result, tuple)
|
||||
assert len(result) >= 3
|
||||
|
||||
|
||||
def test_classify_redemption_error_unclassified() -> None:
|
||||
"""Generic errors are classified as cashu_error with 400 status."""
|
||||
from routstr.wallet import classify_redemption_error
|
||||
|
||||
result = classify_redemption_error(ValueError("unexpected"))
|
||||
# classify_redemption_error classifies all unrecognized errors
|
||||
# as cashu_error with a generic message
|
||||
assert result is not None
|
||||
assert result[0] == "cashu_error"
|
||||
assert result[1] == 400
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Store readiness: store_cashu_transaction succeeds
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_cashu_transaction_succeeds_normally() -> None:
|
||||
"""Normal store_cashu_transaction returns True on success."""
|
||||
from routstr.core.db import store_cashu_transaction
|
||||
|
||||
with patch("routstr.core.db.create_session") as mock_create:
|
||||
mock_session = AsyncMock()
|
||||
mock_session.commit = AsyncMock()
|
||||
mock_session.__aenter__ = AsyncMock(return_value=mock_session)
|
||||
mock_session.__aexit__ = AsyncMock(return_value=None)
|
||||
mock_create.return_value = mock_session
|
||||
|
||||
result = await store_cashu_transaction(
|
||||
token="cashuAtest",
|
||||
amount=1000,
|
||||
unit="sat",
|
||||
typ="in",
|
||||
request_id="req-test",
|
||||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# get_balance
|
||||
# ===========================================================================
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_balance_returns_integer() -> None:
|
||||
"""get_balance returns an integer balance from wallet."""
|
||||
from routstr.wallet import get_balance
|
||||
|
||||
mock_wallet = Mock()
|
||||
mock_wallet.available_balance = Mock(amount=50000)
|
||||
mock_wallet.load_mint = AsyncMock()
|
||||
mock_wallet.load_proofs = AsyncMock()
|
||||
|
||||
with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet):
|
||||
balance = await get_balance("sat")
|
||||
assert isinstance(balance, int)
|
||||
assert balance == 50000
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Periodic task structure verification
|
||||
# ===========================================================================
|
||||
|
||||
def test_periodic_payout_has_loop_and_error_handling() -> None:
|
||||
"""periodic_payout runs in a loop with error handling."""
|
||||
import inspect
|
||||
|
||||
from routstr import wallet
|
||||
|
||||
source = inspect.getsource(wallet.periodic_payout)
|
||||
assert "while True" in source
|
||||
assert "except" in source, "Must have error handling"
|
||||
|
||||
|
||||
def test_periodic_refund_sweep_has_error_handling() -> None:
|
||||
"""Refund sweep catches errors to stay alive."""
|
||||
import inspect
|
||||
|
||||
from routstr import wallet
|
||||
|
||||
source = inspect.getsource(wallet.periodic_refund_sweep)
|
||||
assert "while True" in source
|
||||
assert "except" in source, "Must have error handling"
|
||||
|
||||
|
||||
def test_periodic_routstr_fee_payout_structure() -> None:
|
||||
"""Fee payout loop handles missing LN address gracefully."""
|
||||
import inspect
|
||||
|
||||
from routstr import wallet
|
||||
|
||||
source = inspect.getsource(wallet.periodic_routstr_fee_payout)
|
||||
# Returns early if ROUTSTR_LN_ADDRESS not set
|
||||
assert "ROUTSTR_LN_ADDRESS" in source
|
||||
assert "return" in source or "skip" in source.lower()
|
||||
@@ -0,0 +1,171 @@
|
||||
"""Tests asserting CORRECT behavior for the streaming billing fallback.
|
||||
|
||||
These tests FAIL against current main because the zero-cost fallback at
|
||||
base.py:1012-1030 gives users free service on billing errors.
|
||||
|
||||
Correct behavior required:
|
||||
1. Billing errors must NOT hardcode total_msats=0 — free service is theft
|
||||
2. Reserved balance must be released when billing fails
|
||||
3. Error must be logged at CRITICAL level, not just logger.exception
|
||||
4. The except clause must be narrow, not catch-all Exception
|
||||
"""
|
||||
|
||||
import inspect
|
||||
|
||||
# ===========================================================================
|
||||
# RED TESTS: No hardcoded zero-cost on billing error
|
||||
# ===========================================================================
|
||||
|
||||
def test_billing_error_must_not_hardcode_zero_cost() -> None:
|
||||
"""FIX REQUIRED: billing errors must not result in zero-cost billing.
|
||||
|
||||
base.py:1012-1030 substitutes total_msats=0, total_usd=0.0 when
|
||||
adjust_payment_for_tokens raises ANY exception. This means:
|
||||
- User gets free inference
|
||||
- Reserved balance is never released
|
||||
- Operator has no idea money was lost
|
||||
"""
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
source = inspect.getsource(
|
||||
BaseUpstreamProvider.handle_streaming_chat_completion
|
||||
)
|
||||
|
||||
fallback_start = source.find("Error during usage finalization")
|
||||
assert fallback_start > 0, (
|
||||
"Fallback block exists — it must be removed or fixed"
|
||||
)
|
||||
|
||||
fallback_section = source[fallback_start : fallback_start + 600]
|
||||
|
||||
has_zero_msats = '"total_msats": 0' in fallback_section
|
||||
has_zero_usd = '"total_usd": 0.0' in fallback_section
|
||||
|
||||
assert not has_zero_msats, (
|
||||
"FIX REQUIRED: total_msats is hardcoded to 0 on billing error. "
|
||||
"User gets free service. Fix: propagate the error as a 500 response "
|
||||
"with the token refunded to the user."
|
||||
)
|
||||
assert not has_zero_usd, (
|
||||
"FIX REQUIRED: total_usd is hardcoded to 0.0. No billing occurs. "
|
||||
"Fix: propagate the error."
|
||||
)
|
||||
|
||||
|
||||
def test_billing_error_must_release_reserved_balance() -> None:
|
||||
"""FIX REQUIRED: billing errors must release the reserved balance.
|
||||
|
||||
When adjust_payment_for_tokens fails, the reserved balance on the
|
||||
API key must be released. Currently it's stuck forever.
|
||||
"""
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
source = inspect.getsource(
|
||||
BaseUpstreamProvider.handle_streaming_chat_completion
|
||||
)
|
||||
|
||||
fallback_start = source.find("Error during usage finalization")
|
||||
assert fallback_start > 0, "Fallback exists"
|
||||
|
||||
fallback_section = source[fallback_start : fallback_start + 600]
|
||||
|
||||
has_release = any(
|
||||
kw in fallback_section
|
||||
for kw in ["reserved_balance", "release_reservation", "adjust_reserved",
|
||||
"reset_reserved", "clear_reserved"]
|
||||
)
|
||||
|
||||
assert has_release, (
|
||||
"FIX REQUIRED: Zero-cost fallback does NOT release the reserved "
|
||||
"balance. Funds are permanently stuck. Fix: add reserved_balance "
|
||||
"release in the error path."
|
||||
)
|
||||
|
||||
|
||||
def test_billing_error_catch_is_too_broad() -> None:
|
||||
"""FIX REQUIRED: except clause must not catch all Exception types.
|
||||
|
||||
`except Exception as e:` catches transient DB errors, logic bugs,
|
||||
and serialization failures — all resulting in free service.
|
||||
The catch should be specific (e.g., TemporaryDBError) or the error
|
||||
should propagate as a 500.
|
||||
"""
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
source = inspect.getsource(
|
||||
BaseUpstreamProvider.handle_streaming_chat_completion
|
||||
)
|
||||
|
||||
fallback_start = source.find("Error during usage finalization")
|
||||
# Look at the except clause above the fallback
|
||||
pre_fallback = source[max(0, fallback_start - 250) : fallback_start]
|
||||
|
||||
assert "except Exception" not in pre_fallback, (
|
||||
"FIX REQUIRED: The except clause catches all Exception types. "
|
||||
"A transient DB hiccup results in free inference. "
|
||||
"Fix: narrow the exception type or propagate the error."
|
||||
)
|
||||
|
||||
|
||||
def test_billing_error_must_log_critical() -> None:
|
||||
"""FIX REQUIRED: billing failure must log at CRITICAL level.
|
||||
|
||||
Currently uses logger.exception() which is ERROR level.
|
||||
A billing failure means the operator is losing money — this must
|
||||
be CRITICAL so monitoring/monitoring systems catch it.
|
||||
"""
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
source = inspect.getsource(
|
||||
BaseUpstreamProvider.handle_streaming_chat_completion
|
||||
)
|
||||
|
||||
fallback_start = source.find("Error during usage finalization")
|
||||
fallback_section = source[fallback_start : fallback_start + 600]
|
||||
|
||||
has_critical = "CRITICAL" in fallback_section or "critical" in fallback_section
|
||||
|
||||
assert has_critical, (
|
||||
"FIX REQUIRED: Billing error is logged at ERROR level. "
|
||||
"Money is being lost — this must be CRITICAL so operators "
|
||||
"get alerted."
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# RED TESTS: Messages streaming billing
|
||||
# ===========================================================================
|
||||
|
||||
def test_messages_streaming_no_silent_billing_failure() -> None:
|
||||
"""FIX REQUIRED: messages streaming must not silently swallow billing errors.
|
||||
|
||||
handle_streaming_messages_completion uses `except Exception: pass`
|
||||
for the finalize path, silently dropping the billing attachment.
|
||||
After the fix, this catch block must either:
|
||||
- Log at CRITICAL level with the error details
|
||||
- Propagate the error to surface an HTTP 500
|
||||
- Release reserved balance and refund the token
|
||||
"""
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
source = inspect.getsource(
|
||||
BaseUpstreamProvider.handle_streaming_messages_completion
|
||||
)
|
||||
|
||||
# After fix: the silent pass in finalize_without_usage must be replaced
|
||||
# The fix must include at least one of: CRITICAL logging, error propagation,
|
||||
# or balance release in the error path.
|
||||
|
||||
# The silent pass must NOT exist around billing finalization
|
||||
silent_pass_exists = False
|
||||
for segment in source.split("except Exception:"):
|
||||
if "adjust_payment_for_tokens" in segment:
|
||||
if "pass" in segment[:150]:
|
||||
silent_pass_exists = True
|
||||
break
|
||||
|
||||
assert not silent_pass_exists, (
|
||||
"FIX REQUIRED: finalize_without_usage in messages streaming "
|
||||
"silently swallows billing errors with `except Exception: pass`. "
|
||||
"User gets unbilled inference with no log record."
|
||||
)
|
||||
+3
-3
@@ -4,15 +4,15 @@ FROM base AS deps
|
||||
RUN apk add --no-cache libc6-compat
|
||||
WORKDIR /app
|
||||
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
|
||||
|
||||
COPY package.json pnpm-lock.yaml* ./
|
||||
COPY package.json pnpm-lock.yaml* pnpm-workspace.yaml* ./
|
||||
RUN pnpm install --frozen-lockfile
|
||||
|
||||
FROM base AS builder
|
||||
WORKDIR /app
|
||||
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
|
||||
|
||||
COPY --from=deps /app/node_modules ./node_modules
|
||||
COPY . .
|
||||
|
||||
+3
-3
@@ -6,14 +6,14 @@ RUN apk add --no-cache libc6-compat
|
||||
WORKDIR /app
|
||||
|
||||
# Copy package files
|
||||
COPY package.json pnpm-lock.yaml* ./
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
|
||||
COPY package.json pnpm-lock.yaml* pnpm-workspace.yaml* ./
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
|
||||
RUN pnpm install --frozen-lockfile
|
||||
|
||||
# Build the UI
|
||||
FROM base AS builder
|
||||
WORKDIR /app
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
|
||||
COPY --from=deps /app/node_modules ./node_modules
|
||||
COPY . .
|
||||
|
||||
|
||||
@@ -86,7 +86,7 @@ export function AddModelForm({
|
||||
|
||||
return (
|
||||
<Dialog open={isOpen} onOpenChange={handleClose}>
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[600px]'>
|
||||
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[600px]'>
|
||||
<DialogHeader>
|
||||
<DialogTitle className='flex items-center gap-2'>
|
||||
<Plus className='h-5 w-5' />
|
||||
|
||||
@@ -438,7 +438,7 @@ export function AddProviderModelDialog({
|
||||
|
||||
return (
|
||||
<Dialog open={isOpen} onOpenChange={(open) => !open && onClose()}>
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[720px]'>
|
||||
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[720px]'>
|
||||
<DialogHeader>
|
||||
<DialogTitle className='flex items-center gap-2'>
|
||||
<Plus className='h-4 w-4' />
|
||||
|
||||
@@ -177,7 +177,7 @@ export function CollectModelsDialog({
|
||||
|
||||
return (
|
||||
<Dialog open={isOpen} onOpenChange={handleClose}>
|
||||
<DialogContent className='max-h-[80vh] sm:max-w-[700px]'>
|
||||
<DialogContent className='max-h-[80dvh] sm:max-w-[700px]'>
|
||||
<DialogHeader>
|
||||
<DialogTitle className='flex items-center gap-2'>
|
||||
<Download className='h-5 w-5' />
|
||||
|
||||
@@ -25,7 +25,7 @@ export function CostCalculatorDialog({
|
||||
}: CostCalculatorDialogProps) {
|
||||
return (
|
||||
<Dialog open={isOpen} onOpenChange={onClose}>
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[700px]'>
|
||||
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[700px]'>
|
||||
<DialogHeader>
|
||||
<DialogTitle className='flex items-center gap-2'>
|
||||
<Calculator className='h-5 w-5' />
|
||||
|
||||
@@ -100,7 +100,7 @@ export function EditGroupForm({
|
||||
|
||||
return (
|
||||
<Dialog open={isOpen} onOpenChange={handleClose}>
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[700px]'>
|
||||
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[700px]'>
|
||||
<DialogHeader>
|
||||
<DialogTitle className='flex items-center gap-2'>
|
||||
<Users className='h-5 w-5' />
|
||||
|
||||
@@ -213,7 +213,7 @@ export function ProviderCard({
|
||||
</CardHeader>
|
||||
|
||||
<Dialog open={isKeyModalOpen} onOpenChange={setIsKeyModalOpen}>
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[500px]'>
|
||||
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[500px]'>
|
||||
<DialogHeader>
|
||||
<DialogTitle>
|
||||
{provider.api_key
|
||||
|
||||
@@ -53,7 +53,7 @@ export function ProviderFormDialogContent({
|
||||
availableMints,
|
||||
}: ProviderFormDialogContentProps) {
|
||||
return (
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[500px]'>
|
||||
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[500px]'>
|
||||
<DialogHeader>
|
||||
<DialogTitle>{title}</DialogTitle>
|
||||
<DialogDescription>{description}</DialogDescription>
|
||||
|
||||
@@ -112,7 +112,7 @@ export function ProviderBalance({
|
||||
</Button>
|
||||
|
||||
<Dialog open={isTopupDialogOpen} onOpenChange={handleCloseDialog}>
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-md'>
|
||||
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-md'>
|
||||
<DialogHeader>
|
||||
<DialogTitle>Top Up Balance</DialogTitle>
|
||||
<DialogDescription>
|
||||
|
||||
@@ -239,7 +239,7 @@ export function RoutstrProviderCard({
|
||||
</Card>
|
||||
|
||||
<Dialog open={isKeyDialogOpen} onOpenChange={setIsKeyDialogOpen}>
|
||||
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-lg'>
|
||||
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-lg'>
|
||||
<DialogHeader>
|
||||
<DialogTitle>
|
||||
{hasApiKey ? 'Create New Key on Upstream Node' : 'Create API Key'}
|
||||
|
||||
@@ -116,9 +116,24 @@ export function AdminSettings() {
|
||||
setSaving(true);
|
||||
setError('');
|
||||
|
||||
const updatedData = await AdminService.updateSettings(settings);
|
||||
setSettings(updatedData as SettingsData);
|
||||
setInitialSettings(updatedData as SettingsData);
|
||||
// The nsec is a secret with its own endpoint (the general settings PATCH
|
||||
// strips it); only send it when the operator actually changed it, so an
|
||||
// untouched redacted value is never written back. The npub is derived
|
||||
// server-side from the new nsec — fold it into the settings payload so the
|
||||
// persisted blob stays consistent with the stored key.
|
||||
let settingsPayload = settings;
|
||||
if (hasFieldChanged('nsec')) {
|
||||
const result = await AdminService.updateNsec(
|
||||
(settings.nsec as string) || ''
|
||||
);
|
||||
settingsPayload = { ...settings, npub: result.npub };
|
||||
}
|
||||
|
||||
const updatedData = (await AdminService.updateSettings(
|
||||
settingsPayload
|
||||
)) as SettingsData;
|
||||
setSettings(updatedData);
|
||||
setInitialSettings(updatedData);
|
||||
toast.success('Settings saved successfully');
|
||||
} catch (err) {
|
||||
const message =
|
||||
@@ -140,8 +155,8 @@ export function AdminSettings() {
|
||||
return;
|
||||
}
|
||||
|
||||
if (passwordData.new_password.length < 6) {
|
||||
setPasswordError('New password must be at least 6 characters');
|
||||
if (passwordData.new_password.length < 8) {
|
||||
setPasswordError('New password must be at least 8 characters');
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -944,7 +959,7 @@ export function AdminSettings() {
|
||||
new_password: e.target.value,
|
||||
}))
|
||||
}
|
||||
placeholder='Enter new password (min 6 characters)'
|
||||
placeholder='Enter new password (min 8 characters)'
|
||||
/>
|
||||
</div>
|
||||
|
||||
|
||||
@@ -626,7 +626,7 @@ export function TopModelsUsageChart({
|
||||
className={cn(
|
||||
'aspect-auto w-full',
|
||||
isFullscreen
|
||||
? 'h-[calc(100vh-220px)] min-h-[340px] sm:h-[calc(100vh-260px)] sm:min-h-[420px]'
|
||||
? 'h-[calc(100dvh-220px)] min-h-[340px] sm:h-[calc(100dvh-260px)] sm:min-h-[420px]'
|
||||
: 'h-[260px] sm:h-[340px]'
|
||||
)}
|
||||
onMouseLeave={() => {
|
||||
|
||||
@@ -130,7 +130,7 @@ function DialogContent({
|
||||
<DrawerPrimitive.Content
|
||||
data-slot='dialog-content'
|
||||
className={cn(
|
||||
'border-border bg-card data-closed:animate-out data-open:animate-in data-closed:slide-out-to-bottom data-open:slide-in-from-bottom fixed inset-x-0 bottom-0 z-50 grid max-h-[90vh] w-full gap-4 rounded-t-xl border-t p-4 pb-[calc(1rem+env(safe-area-inset-bottom))] shadow-lg',
|
||||
'border-border bg-card data-closed:animate-out data-open:animate-in data-closed:slide-out-to-bottom data-open:slide-in-from-bottom fixed inset-x-0 bottom-0 z-50 grid max-h-[90dvh] w-full gap-4 rounded-t-xl border-t p-4 pb-[calc(1rem+env(safe-area-inset-bottom))] shadow-lg',
|
||||
className
|
||||
)}
|
||||
{...(props as React.ComponentProps<typeof DrawerPrimitive.Content>)}
|
||||
|
||||
@@ -56,7 +56,7 @@ function DrawerContent({
|
||||
<DrawerPrimitive.Content
|
||||
data-slot='drawer-content'
|
||||
className={cn(
|
||||
'bg-background group/drawer-content fixed z-50 flex h-auto flex-col text-sm data-[vaul-drawer-direction=bottom]:inset-x-0 data-[vaul-drawer-direction=bottom]:bottom-0 data-[vaul-drawer-direction=bottom]:mt-24 data-[vaul-drawer-direction=bottom]:max-h-[80vh] data-[vaul-drawer-direction=bottom]:rounded-t-xl data-[vaul-drawer-direction=bottom]:border-t data-[vaul-drawer-direction=bottom]:pb-[calc(0.5rem+env(safe-area-inset-bottom))] data-[vaul-drawer-direction=left]:inset-y-0 data-[vaul-drawer-direction=left]:left-0 data-[vaul-drawer-direction=left]:w-3/4 data-[vaul-drawer-direction=left]:rounded-r-xl data-[vaul-drawer-direction=left]:border-r data-[vaul-drawer-direction=right]:inset-y-0 data-[vaul-drawer-direction=right]:right-0 data-[vaul-drawer-direction=right]:w-3/4 data-[vaul-drawer-direction=right]:rounded-l-xl data-[vaul-drawer-direction=right]:border-l data-[vaul-drawer-direction=top]:inset-x-0 data-[vaul-drawer-direction=top]:top-0 data-[vaul-drawer-direction=top]:mb-24 data-[vaul-drawer-direction=top]:max-h-[80vh] data-[vaul-drawer-direction=top]:rounded-b-xl data-[vaul-drawer-direction=top]:border-b data-[vaul-drawer-direction=left]:sm:max-w-sm data-[vaul-drawer-direction=right]:sm:max-w-sm',
|
||||
'bg-background group/drawer-content fixed z-50 flex h-auto flex-col text-sm data-[vaul-drawer-direction=bottom]:inset-x-0 data-[vaul-drawer-direction=bottom]:bottom-0 data-[vaul-drawer-direction=bottom]:mt-24 data-[vaul-drawer-direction=bottom]:max-h-[80dvh] data-[vaul-drawer-direction=bottom]:rounded-t-xl data-[vaul-drawer-direction=bottom]:border-t data-[vaul-drawer-direction=bottom]:pb-[calc(0.5rem+env(safe-area-inset-bottom))] data-[vaul-drawer-direction=left]:inset-y-0 data-[vaul-drawer-direction=left]:left-0 data-[vaul-drawer-direction=left]:w-3/4 data-[vaul-drawer-direction=left]:rounded-r-xl data-[vaul-drawer-direction=left]:border-r data-[vaul-drawer-direction=right]:inset-y-0 data-[vaul-drawer-direction=right]:right-0 data-[vaul-drawer-direction=right]:w-3/4 data-[vaul-drawer-direction=right]:rounded-l-xl data-[vaul-drawer-direction=right]:border-l data-[vaul-drawer-direction=top]:inset-x-0 data-[vaul-drawer-direction=top]:top-0 data-[vaul-drawer-direction=top]:mb-24 data-[vaul-drawer-direction=top]:max-h-[80dvh] data-[vaul-drawer-direction=top]:rounded-b-xl data-[vaul-drawer-direction=top]:border-b data-[vaul-drawer-direction=left]:sm:max-w-sm data-[vaul-drawer-direction=right]:sm:max-w-sm',
|
||||
className
|
||||
)}
|
||||
{...props}
|
||||
|
||||
@@ -291,7 +291,7 @@ export function UsageMetricsChart({
|
||||
className={cn(
|
||||
'aspect-auto w-full',
|
||||
isFullscreen
|
||||
? 'h-[calc(100vh-220px)] min-h-[340px] sm:h-[calc(100vh-260px)] sm:min-h-[420px]'
|
||||
? 'h-[calc(100dvh-220px)] min-h-[340px] sm:h-[calc(100dvh-260px)] sm:min-h-[420px]'
|
||||
: 'h-[260px] sm:h-[340px]'
|
||||
)}
|
||||
config={chartConfig}
|
||||
|
||||
@@ -801,6 +801,15 @@ export class AdminService {
|
||||
);
|
||||
}
|
||||
|
||||
static async updateNsec(
|
||||
nsec: string
|
||||
): Promise<{ ok: boolean; npub: string }> {
|
||||
return await apiClient.patch<{ ok: boolean; npub: string }>(
|
||||
'/admin/api/nsec',
|
||||
{ nsec }
|
||||
);
|
||||
}
|
||||
|
||||
static async login(password: string): Promise<{
|
||||
ok: boolean;
|
||||
token: string;
|
||||
|
||||
+3
-9
@@ -45,7 +45,7 @@
|
||||
"@radix-ui/react-tooltip": "^1.2.8",
|
||||
"@tanstack/react-query": "^5.90.21",
|
||||
"@tanstack/react-table": "^8.21.3",
|
||||
"axios": "^1.13.5",
|
||||
"axios": "^1.16.0",
|
||||
"class-variance-authority": "^0.7.1",
|
||||
"clsx": "^2.1.1",
|
||||
"cmdk": "^1.1.1",
|
||||
@@ -54,7 +54,7 @@
|
||||
"geist": "^1.7.0",
|
||||
"input-otp": "^1.4.2",
|
||||
"lucide-react": "^0.575.0",
|
||||
"next": "16.1.6",
|
||||
"next": "16.2.6",
|
||||
"next-themes": "^0.4.6",
|
||||
"qrcode": "^1.5.4",
|
||||
"qrcode.react": "^4.2.0",
|
||||
@@ -81,7 +81,7 @@
|
||||
"@types/react": "^19.2.14",
|
||||
"@types/react-dom": "^19.2.3",
|
||||
"eslint": "^9.7.0",
|
||||
"eslint-config-next": "16.1.6",
|
||||
"eslint-config-next": "16.2.6",
|
||||
"eslint-config-prettier": "^10.1.8",
|
||||
"eslint-plugin-prettier": "^5.5.5",
|
||||
"eslint-plugin-react": "^7.37.5",
|
||||
@@ -89,11 +89,5 @@
|
||||
"prettier-plugin-tailwindcss": "^0.7.2",
|
||||
"tailwindcss": "^4.2.0",
|
||||
"typescript": "^5.9.3"
|
||||
},
|
||||
"pnpm": {
|
||||
"onlyBuiltDependencies": [
|
||||
"sharp",
|
||||
"unrs-resolver"
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user