mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-04 17:14:38 +00:00
Compare commits
230
Commits
v0.4.4.e2ee
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c971862ac6 | ||
|
|
22b35ff93d | ||
|
|
2f2820eb33 | ||
|
|
a6c129c02d | ||
|
|
a2e2a5c662 | ||
|
|
8e8a9a46b6 | ||
|
|
47d0d87a88 | ||
|
|
3b2c5a0671 | ||
|
|
0609c5ed77 | ||
|
|
694bc04623 | ||
|
|
da859f2f84 | ||
|
|
dd8c4a9a8a | ||
|
|
2b4e4c2430 | ||
|
|
8c3d8f52ba | ||
|
|
e903aa3a9f | ||
|
|
f3eefc2638 | ||
|
|
3463149b38 | ||
|
|
dc13cde00c | ||
|
|
55dba5136c | ||
|
|
c0aad3b3ab | ||
|
|
19236ecc9d | ||
|
|
7a2b485af6 | ||
|
|
2ec6b27200 | ||
|
|
f9980e5c66 | ||
|
|
98fe5a37cd | ||
|
|
60566313dc | ||
|
|
443c910b9e | ||
|
|
88d301398b | ||
|
|
4cc9aef61f | ||
|
|
3befe063f4 | ||
|
|
f8adaee362 | ||
|
|
895ea90bfa | ||
|
|
c4d27ba02a | ||
|
|
5ea5024608 | ||
|
|
bb2a05b67c | ||
|
|
c75dee147a | ||
|
|
e2f89a2645 | ||
|
|
48c11eb7bc | ||
|
|
c829685f80 | ||
|
|
16fc548b48 | ||
|
|
c5da73f1e9 | ||
|
|
06dba681c5 | ||
|
|
f96acbb99c | ||
|
|
0a00527626 | ||
|
|
73a3f12469 | ||
|
|
39f801561b | ||
|
|
ff55788e2d | ||
|
|
e38cd32fa3 | ||
|
|
423e2cba73 | ||
|
|
1138cdd4ef | ||
|
|
f15eab9f10 | ||
|
|
344c3c5f21 | ||
|
|
4c6bc49e07 | ||
|
|
1d4b8d7cb2 | ||
|
|
1131c2d583 | ||
|
|
7108d554c8 | ||
|
|
1eddf89d52 | ||
|
|
b0c70ecddc | ||
|
|
a2cedd6769 | ||
|
|
7b0ade3987 | ||
|
|
81c0ff57e9 | ||
|
|
2410a4a6ce | ||
|
|
1b09639265 | ||
|
|
7394b10e75 | ||
|
|
4c292580e8 | ||
|
|
04879cad56 | ||
|
|
ab80657507 | ||
|
|
b94d95fc2b | ||
|
|
27f81dbf42 | ||
|
|
66ba31d0df | ||
|
|
16679c1f4e | ||
|
|
a9593ab416 | ||
|
|
6221ee8152 | ||
|
|
65ea28cb85 | ||
|
|
2ed20b1b85 | ||
|
|
88fe9758a3 | ||
|
|
03b00f3eb6 | ||
|
|
7e8c2033ae | ||
|
|
52742a6a04 | ||
|
|
030d8b61ce | ||
|
|
f47e16aa61 | ||
|
|
eb6a612e5a | ||
|
|
c2a2d76eae | ||
|
|
d4657dca5a | ||
|
|
7106dfe330 | ||
|
|
1045f1061b | ||
|
|
0415806c5a | ||
|
|
5d5c849180 | ||
|
|
b8700dde40 | ||
|
|
4872c318d5 | ||
|
|
56a67c0a86 | ||
|
|
afcb3f7cda | ||
|
|
14748c28fb | ||
|
|
8aa2fa5c4a | ||
|
|
a024d5be5e | ||
|
|
c40723c1e2 | ||
|
|
6c5c3f0b5f | ||
|
|
86adee9b1b | ||
|
|
92246b78d0 | ||
|
|
6023c03959 | ||
|
|
dbe7a53afd | ||
|
|
e0c74e3a46 | ||
|
|
18965b4ea4 | ||
|
|
97dc10a8ad | ||
|
|
a9a6381614 | ||
|
|
040799a4d7 | ||
|
|
87850c97b9 | ||
|
|
fcc87718ff | ||
|
|
50437a1cc6 | ||
|
|
7f5a0cf1ae | ||
|
|
3f850c54f4 | ||
|
|
df4d4c44e6 | ||
|
|
02cd2cfeec | ||
|
|
0aebfc6dbe | ||
|
|
b76fa17f81 | ||
|
|
ed8ac914f9 | ||
|
|
6c762f0c3d | ||
|
|
2b4f70442a | ||
|
|
99724cc5f5 | ||
|
|
e19f679609 | ||
|
|
bd1edcef26 | ||
|
|
586af15a1b | ||
|
|
3e906605a0 | ||
|
|
22a94a68a4 | ||
|
|
4defe4f227 | ||
|
|
f8125a8a2d | ||
|
|
999a5634fa | ||
|
|
2c218cce49 | ||
|
|
fa0b366f9a | ||
|
|
a3a4d69ed3 | ||
|
|
c9533c872a | ||
|
|
90da3803c6 | ||
|
|
8b3b59e176 | ||
|
|
be5e323c68 | ||
|
|
f6d1a41728 | ||
|
|
b3bf1f0e90 | ||
|
|
3035292d2a | ||
|
|
ac436d0862 | ||
|
|
19082231f9 | ||
|
|
281607108c | ||
|
|
8314f3c1b0 | ||
|
|
1a9041766b | ||
|
|
01c01fe8ad | ||
|
|
892aed61cc | ||
|
|
1957e716a3 | ||
|
|
69f19ff991 | ||
|
|
8b942f3c14 | ||
|
|
09e1c7bf2d | ||
|
|
cc2a96e2ef | ||
|
|
c6733dbb62 | ||
|
|
93ab1d927b | ||
|
|
39970d8bee | ||
|
|
d7c401d204 | ||
|
|
d44b98fd0d | ||
|
|
6fa3610423 | ||
|
|
65702171e4 | ||
|
|
65abcbce92 | ||
|
|
57fef44ce7 | ||
|
|
c671d277d3 | ||
|
|
9460e24f21 | ||
|
|
f31929b538 | ||
|
|
129ea7bb76 | ||
|
|
b88c6d84fa | ||
|
|
40153d4c36 | ||
|
|
238beb3e52 | ||
|
|
40bf976fbc | ||
|
|
eae20f04a7 | ||
|
|
dc833d9f5f | ||
|
|
31e7ff8fbd | ||
|
|
a1850dfa55 | ||
|
|
9eacabde0e | ||
|
|
73f52ef1fd | ||
|
|
110ec5da5d | ||
|
|
a5cee796c9 | ||
|
|
aabf08a930 | ||
|
|
07c15de106 | ||
|
|
aea24d23da | ||
|
|
e100c77288 | ||
|
|
057a752b1e | ||
|
|
5daa2602f5 | ||
|
|
acb630f6cf | ||
|
|
1230d528de | ||
|
|
d80912b10e | ||
|
|
8f087bf07a | ||
|
|
d23c90b939 | ||
|
|
d8db2a3051 | ||
|
|
0bbbf902cd | ||
|
|
7ed18a9d02 | ||
|
|
744321153f | ||
|
|
b81c5add6a | ||
|
|
02109616b9 | ||
|
|
b30346bcb6 | ||
|
|
956a1ac3e1 | ||
|
|
4287f038cf | ||
|
|
82fd2c08a7 | ||
|
|
a1118a510d | ||
|
|
27dedfdf77 | ||
|
|
ebb2267844 | ||
|
|
7c2e2ce512 | ||
|
|
89b93fde98 | ||
|
|
c48693973b | ||
|
|
407f4745a0 | ||
|
|
806771df13 | ||
|
|
0f7d3d2f86 | ||
|
|
3c1b6d5f17 | ||
|
|
a9d5a52b32 | ||
|
|
af2abc3a1c | ||
|
|
b36bb56a30 | ||
|
|
9b4b4d734f | ||
|
|
97172641f8 | ||
|
|
46bbba7027 | ||
|
|
edfb0f127c | ||
|
|
65ccf83dbd | ||
|
|
2baea4149e | ||
|
|
64d9460711 | ||
|
|
dc25659cff | ||
|
|
349d8dd009 | ||
|
|
6097193442 | ||
|
|
e2aa02307b | ||
|
|
66bc1260a4 | ||
|
|
9dc4c979b8 | ||
|
|
ec0fd1143b | ||
|
|
ae9748db02 | ||
|
|
ccee76b31e | ||
|
|
9f04525c82 | ||
|
|
2a27fb239e | ||
|
|
ffde661d93 | ||
|
|
7f49ba1771 | ||
|
|
949dc433f1 | ||
|
|
b509581950 |
+32
-2
@@ -5,21 +5,51 @@ UPSTREAM_API_KEY=your-upstream-api-key
|
||||
# Tinfoil (confidential inference enclaves, EHBP)
|
||||
# TINFOIL_API_KEY=your-tinfoil-api-key
|
||||
|
||||
# ADMIN_PASSWORD=secure-admin-password
|
||||
# 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
|
||||
# Pool controls are validated at boot, sourced only from the environment, and
|
||||
# logged at startup. Keep total capacity across all workers below the database
|
||||
# connection limit. Pre-ping is automatic for networked backends; SQLite may
|
||||
# explicitly opt in if desired.
|
||||
# DATABASE_POOL_SIZE=10
|
||||
# DATABASE_MAX_OVERFLOW=20
|
||||
# DATABASE_POOL_TIMEOUT=15
|
||||
# DATABASE_POOL_RECYCLE=1800
|
||||
# DATABASE_POOL_PRE_PING=false
|
||||
# Warn when a checkout is held this many seconds.
|
||||
# DATABASE_POOL_HOLD_WARN_SECONDS=10
|
||||
# SQLite serialises writes; increasing its pool can trade pool timeouts for
|
||||
# "database is locked" errors rather than increasing write throughput.
|
||||
|
||||
# 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"
|
||||
# ENABLE_ANALYTICS_SHARING=true
|
||||
# CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me"
|
||||
# MINT_OPERATION_CONCURRENCY=4
|
||||
# MINT_OPERATION_TIMEOUT_SECONDS=30
|
||||
# MINT_MAX_CONCURRENCY=4
|
||||
# MINT_RETRY_MAX_ATTEMPTS=3
|
||||
# RECEIVE_LN_ADDRESS=
|
||||
# REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS=900
|
||||
|
||||
# Custom Pricing Configuration
|
||||
# MODEL_BASED_PRICING=true
|
||||
|
||||
@@ -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
|
||||
|
||||
+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:
|
||||
|
||||
@@ -327,6 +327,62 @@ GET /v1/models
|
||||
}
|
||||
```
|
||||
|
||||
### List Model Paths
|
||||
|
||||
Get the selectable upstream routes for each advertised model. This endpoint is
|
||||
discovery-only; request-side selection will be added separately.
|
||||
|
||||
```http
|
||||
GET /v1/models/paths
|
||||
```
|
||||
|
||||
**Response:**
|
||||
|
||||
```json
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"id": "anthropic/claude-sonnet-4",
|
||||
"paths": [
|
||||
{
|
||||
"path": "url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=anthropic%2Fclaude-sonnet-4",
|
||||
"provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"},
|
||||
"endpoint": null
|
||||
},
|
||||
{
|
||||
"path": "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4&endpoint=google-vertex%2Fus",
|
||||
"provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"},
|
||||
"endpoint": {"tag": "google-vertex/us", "name": "Google"}
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"updated_at": 1753500000
|
||||
}
|
||||
```
|
||||
|
||||
`path` is an opaque, percent-encoded selector. Clients must store and return it
|
||||
unchanged rather than parsing or reconstructing it. It identifies the exact
|
||||
configured route with `url`, `provider-id`, and `model-id`. To avoid exposing
|
||||
private network details, a configured private IP address or any URL with an
|
||||
explicit port is advertised as `http://localhost`. OpenRouter routes additionally
|
||||
preserve the exact machine-readable endpoint `tag`. Provider slugs/types and
|
||||
endpoint names remain display data. When request-side selection is implemented,
|
||||
an endpoint tag must not silently fall back to another backend.
|
||||
|
||||
### List Paths for One Model
|
||||
|
||||
Use the exact model ID advertised by `/v1/models`. The query parameter safely
|
||||
supports IDs containing `/`.
|
||||
|
||||
```http
|
||||
GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4
|
||||
```
|
||||
|
||||
The response uses the same path objects and `updated_at` field as the collection
|
||||
endpoint. An unknown model returns `404 Model not found`. A known model whose
|
||||
paths have not been discovered yet returns `200` with an empty `data` array.
|
||||
|
||||
## Wallet Management
|
||||
|
||||
### Create Wallet (Coming Soon)
|
||||
|
||||
+87
-20
@@ -113,42 +113,109 @@ All errors follow a consistent JSON structure:
|
||||
**Status:** 402
|
||||
**Resolution:** Top up API key balance
|
||||
|
||||
#### Invalid Token
|
||||
### Cashu Token Redemption Errors
|
||||
|
||||
These errors are returned when a Cashu token you pay with cannot be redeemed.
|
||||
They apply to every endpoint that accepts a token:
|
||||
|
||||
- **Per-request payment** via the `X-Cashu` header (chat completions + Responses API).
|
||||
- **API key top-up** via `POST /v1/wallet/topup`.
|
||||
- **Minting an API key** from a token sent in `Authorization: Bearer <cashu-token>`.
|
||||
|
||||
All three share one classifier, so the same failure yields the same HTTP status
|
||||
and sanitized message everywhere. Structured error envelopes (`X-Cashu` and
|
||||
`Authorization: Bearer <cashu-token>`) also expose the same `type` and `code` —
|
||||
branch on `type` (or `code` for finer granularity). `POST /v1/wallet/topup`
|
||||
keeps its existing plain-string `detail` envelope, so branch on status there.
|
||||
|
||||
| `type` | Status | `code` | Retryable | Meaning |
|
||||
|--------|--------|--------|-----------|---------|
|
||||
| `token_already_spent` | 400 | `cashu_token_already_spent` | No | The token was already redeemed. |
|
||||
| `invalid_token` | 400 | `invalid_cashu_token` | No | The token is malformed or cannot be decoded. |
|
||||
| `mint_error` | 422 | `cashu_token_swap_fees_exceed_amount` | No | Token value is too small to cover the mint's swap/melt fees. |
|
||||
| `mint_error` | 422 | `cashu_foreign_mint_swap_failed` | No | Swapping the token from a foreign mint to the primary mint failed. |
|
||||
| `mint_unreachable` | 503 | `cashu_mint_unreachable` | **Yes** | The mint could not be reached (DNS failure, refused/reset connection, timeout). The token is fine — retry once the mint recovers. |
|
||||
| `cashu_error` | 400 | `cashu_token_redemption_failed` | No | The token could not be redeemed for another expected reason. |
|
||||
| `cashu_error` | 400 | `cashu_token_zero_value` | No | The token redeemed to zero (empty/dust token, or value fully consumed by fees). |
|
||||
| `token_consumed` | 500 | `cashu_token_consumed` | No | The token was **spent** (melted/redeemed) but crediting it then failed. Do not retry — the token is gone; contact support to reconcile. |
|
||||
| `api_error` | 500 | `internal_error` | Maybe | Unexpected server-side fault during redemption. |
|
||||
|
||||
!!! important "Retry only `mint_unreachable`"
|
||||
Only `mint_unreachable` (503) means the same token will work again later —
|
||||
everything else is a permanent property of the token and must not be
|
||||
blindly retried. Use exponential backoff for the 503. In particular, a
|
||||
`token_consumed` 500 means the mint already spent the token, so a retry
|
||||
would fail as `token_already_spent`.
|
||||
|
||||
#### Mint Unreachable (retryable)
|
||||
|
||||
```json
|
||||
{
|
||||
"error": {
|
||||
"type": "payment_error",
|
||||
"message": "Invalid Cashu token",
|
||||
"code": "invalid_token",
|
||||
"details": {
|
||||
"reason": "Token already spent"
|
||||
}
|
||||
"type": "mint_unreachable",
|
||||
"message": "Cashu mint is unreachable",
|
||||
"code": "cashu_mint_unreachable"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Status:** 400
|
||||
**Resolution:** Use a valid, unspent token
|
||||
**Status:** 503
|
||||
|
||||
#### Mint Unavailable
|
||||
**Resolution:** The token is valid — the mint is temporarily down. Retry with
|
||||
backoff, or pay with a token from a different mint.
|
||||
|
||||
#### Token Already Spent
|
||||
|
||||
```json
|
||||
{
|
||||
"error": {
|
||||
"type": "payment_error",
|
||||
"message": "Cannot connect to Cashu mint",
|
||||
"code": "mint_unavailable",
|
||||
"details": {
|
||||
"mint_url": "https://mint.example.com",
|
||||
"retry_after": 60
|
||||
}
|
||||
"type": "token_already_spent",
|
||||
"message": "Cashu token already spent",
|
||||
"code": "cashu_token_already_spent"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Status:** 503
|
||||
**Resolution:** Try again later or use different mint
|
||||
**Status:** 400
|
||||
|
||||
**Resolution:** Use a fresh, unspent token. Do not retry with the same token.
|
||||
|
||||
#### Response envelope differs by endpoint
|
||||
|
||||
The `error` object above is identical everywhere, but the surrounding envelope
|
||||
depends on how you paid:
|
||||
|
||||
- **`X-Cashu` header payments** (chat + Responses API) return the object at the
|
||||
top level, alongside a `request_id`:
|
||||
|
||||
```json
|
||||
{
|
||||
"error": { "type": "mint_unreachable", "message": "Cashu mint is unreachable", "code": "cashu_mint_unreachable" },
|
||||
"request_id": "req-abc123"
|
||||
}
|
||||
```
|
||||
|
||||
The original token is echoed back in the `X-Cashu` **response header only when
|
||||
it is still spendable** (e.g. `mint_unreachable`, `invalid_cashu_token`, fee
|
||||
errors) so you can recover/retry it. It is **not** echoed for spent/consumed
|
||||
tokens (`cashu_token_already_spent`, `cashu_token_consumed`,
|
||||
`cashu_token_zero_value`, `internal_error`) — retrying those can never succeed.
|
||||
|
||||
- **`Authorization: Bearer <cashu-token>`** (API key minting) wraps it in
|
||||
FastAPI's `detail` field:
|
||||
|
||||
```json
|
||||
{ "detail": { "error": { "type": "mint_unreachable", "message": "Cashu mint is unreachable", "code": "cashu_mint_unreachable" } } }
|
||||
```
|
||||
|
||||
- **`POST /v1/wallet/topup`** returns a plain string message under `detail` —
|
||||
it carries the shared HTTP **status** and **message** (e.g. `503` for an
|
||||
unreachable mint) but not the structured `type`/`code`, so branch on the
|
||||
status code here:
|
||||
|
||||
```json
|
||||
{ "detail": "Cashu mint is unreachable" }
|
||||
```
|
||||
|
||||
### Validation Errors
|
||||
|
||||
@@ -347,7 +414,7 @@ class ErrorHandler:
|
||||
'rate_limit',
|
||||
'upstream_timeout',
|
||||
'model_overloaded',
|
||||
'mint_unavailable'
|
||||
'cashu_mint_unreachable'
|
||||
}
|
||||
|
||||
# Errors requiring user action
|
||||
|
||||
@@ -100,6 +100,7 @@ All errors follow a consistent format:
|
||||
Standard OpenAI-compatible endpoints:
|
||||
|
||||
- **Models**: `/v1/models`
|
||||
- **Model paths**: `/v1/models/paths`, `/v1/models/paths/model?model_id=...`
|
||||
- **Responses**: `/v1/responses`
|
||||
- **Chat Completions**: `/v1/chat/completions`
|
||||
- **Embeddings**: `/v1/embeddings`
|
||||
@@ -302,7 +303,7 @@ Get node metadata:
|
||||
GET /v1/info
|
||||
```
|
||||
|
||||
Supported models and pricing are available at `/v1/models`.
|
||||
Supported models and pricing are available at `/v1/models`. Upstream provider path discovery is available at `/v1/models/paths` and `/v1/models/paths/model?model_id=...`.
|
||||
|
||||
## Next Steps
|
||||
|
||||
|
||||
+28
-36
@@ -53,25 +53,22 @@ The actual EHBP forwarding logic does **not** live in `base.py`.
|
||||
|
||||
### `routstr/upstream/ehbp.py`
|
||||
|
||||
Contains the shared opaque EHBP transport helpers:
|
||||
Contains the shared opaque EHBP transport and billing helpers:
|
||||
|
||||
- `EHBPForwardingTarget` — provider-specific target URL plus extra headers
|
||||
- `forward_ehbp_request()` — forwards the raw encrypted body to an EHBP-capable
|
||||
provider, streams the encrypted response back untouched, and finalizes bearer
|
||||
billing at max cost because usage is encrypted
|
||||
- `forward_ehbp_x_cashu_request()` — redeems the Cashu token, forwards raw,
|
||||
refunds the full token on upstream failure, and refunds any value above
|
||||
`max_cost_for_model` on success
|
||||
- `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)
|
||||
|
||||
### `routstr/upstream/ppqai.py`
|
||||
### Provider support
|
||||
|
||||
- Sets `supports_ehbp = True`.
|
||||
- Implements `get_ehbp_forwarding_target()` to forward to
|
||||
`https://api.ppq.ai/private/v1/...` — the PPQ.AI enclave endpoint that
|
||||
understands EHBP and returns the `Ehbp-Response-Nonce` header.
|
||||
- Adds `X-Private-Model` with the model's `forwarded_model_id` (e.g.
|
||||
`private/kimi-k2-6`). PPQ.AI's billing layer needs this since it can't
|
||||
decrypt the body.
|
||||
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
|
||||
|
||||
@@ -84,10 +81,10 @@ The proxy is a **blind relay** for EHBP requests. It cannot decrypt the body
|
||||
4. Pass through EHBP protocol headers (`Ehbp-Encapsulated-Key` on request,
|
||||
`Ehbp-Response-Nonce` on response)
|
||||
|
||||
Cost tracking happens at the proxy level using `max_cost_for_model` from the
|
||||
model registry. Because EHBP responses are encrypted, Routstr cannot reconcile
|
||||
against token usage. Bearer requests reserve and then finalize max-cost billing;
|
||||
X-Cashu requests redeem the token and refund any amount above max cost.
|
||||
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
|
||||
|
||||
@@ -125,7 +122,7 @@ 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 |
|
||||
| PPQ.AI billing | `X-Private-Model` header | `private/kimi-k2-6` | Proxy sends `forwarded_model_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
|
||||
@@ -135,24 +132,19 @@ 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).
|
||||
- Override the forwarding URL with `X-Tinfoil-Enclave-Url` when the SDK sends it.
|
||||
- Finalize bearer billing with actual token cost via `adjust_payment_for_tokens`.
|
||||
- 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.
|
||||
|
||||
## Not yet tested
|
||||
## Verification status
|
||||
|
||||
These changes were written without integration testing due to the complexity
|
||||
of the full stack (SDK + proxy + PPQ.AI enclave + Cashu mint). Needs end-to-end
|
||||
verification with a real `tinfoil-*` model request.
|
||||
|
||||
Important assumptions to verify:
|
||||
|
||||
- PPQ.AI accepts `/private/v1/...` with `X-Private-Model`.
|
||||
- PPQ.AI enforces consistency between `X-Private-Model` and the encrypted
|
||||
`body.model`, otherwise a malicious client could understate
|
||||
`X-Routstr-Model` for billing.
|
||||
- SDK behavior on non-2xx proxy-generated errors that do not carry
|
||||
`Ehbp-Response-Nonce`.
|
||||
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.
|
||||
|
||||
---
|
||||
|
||||
@@ -45,6 +48,68 @@ Connect to your AI provider(s):
|
||||
| **Upstream URL** | API endpoint (e.g., `https://api.openai.com/v1`) |
|
||||
| **API Key** | Your provider's API key |
|
||||
|
||||
### PPQ Auto Top-up
|
||||
|
||||
PPQ providers can automatically purchase more credits when their USD balance
|
||||
falls below a configured threshold. Configure this per provider in the Admin
|
||||
Dashboard by editing a **PPQ.AI** provider and opening **PPQ Auto Top-up**.
|
||||
There are no environment variables for this feature.
|
||||
|
||||
#### Requirements
|
||||
|
||||
Before enabling auto top-up, make sure that:
|
||||
|
||||
- the PPQ provider has a valid API key;
|
||||
- at least one trusted Cashu mint is configured;
|
||||
- the node wallet has enough **node-owned** funds at one mint to pay the
|
||||
Lightning invoice; client balances are never used; and
|
||||
- the node has a current BTC/USD price for validating the invoice amount.
|
||||
|
||||
| Setting | Description |
|
||||
| ------- | ----------- |
|
||||
| **Enable Auto Top-up** | Enables automatic PPQ credit purchases for this provider. |
|
||||
| **When credits are below (USD)** | Starts a top-up when the reported PPQ balance is below this positive USD value. |
|
||||
| **Purchase this amount (USD)** | Amount of PPQ credit to buy per top-up. Must be a whole number from **1 to 500 USD**. |
|
||||
|
||||
For example, a threshold of `5` and purchase amount of `20` buys 20 USD of
|
||||
credit when the PPQ balance drops below 5 USD.
|
||||
|
||||
#### How it works
|
||||
|
||||
The worker checks eligible providers approximately once per minute. When the
|
||||
balance is below the threshold, it:
|
||||
|
||||
1. verifies the node has enough owner funds before creating an invoice;
|
||||
2. requests a USD-denominated Lightning top-up invoice from PPQ;
|
||||
3. rejects expired, mismatched, or unexpectedly expensive invoices (more than
|
||||
10% above the local BTC/USD estimate);
|
||||
4. pays from the configured Cashu mint with sufficient owner funds; and
|
||||
5. waits for PPQ to confirm that the credit settled.
|
||||
|
||||
Only one attempt can be active for a provider. An attempt that was active at
|
||||
the start of a cycle suppresses another top-up for that entire cycle, even if
|
||||
PPQ reports it settled immediately. This prevents a temporarily stale PPQ
|
||||
balance from causing a duplicate purchase.
|
||||
|
||||
Completed PPQ payments appear in the dashboard transaction history with source
|
||||
`ppq_auto_topup`. The payment record is separate from the internal claim used
|
||||
to prevent concurrent attempts.
|
||||
|
||||
#### Payment recovery
|
||||
|
||||
If the Cashu mint paid the invoice but PPQ settlement cannot be confirmed, the
|
||||
provider card shows **Auto top-up needs review**. A payment still owned by a
|
||||
running worker is shown as **Paying invoice** and cannot be released.
|
||||
|
||||
Before choosing **Release top-up**, manually verify both PPQ and the Cashu mint.
|
||||
Release the claim only when the previous Lightning payment is definitively
|
||||
unable to settle. Releasing an ambiguous payment allows the next cycle to try
|
||||
again and can therefore cause a duplicate top-up.
|
||||
|
||||
Disabling auto top-up prevents new purchases, but the node continues to
|
||||
reconcile an already active payment until it reaches a safe terminal state or
|
||||
requires operator review.
|
||||
|
||||
### Node Identity
|
||||
|
||||
How your node appears to clients:
|
||||
@@ -123,25 +188,52 @@ 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` |
|
||||
| `MINT_OPERATION_CONCURRENCY` | Concurrent mint/unit balance reads | `4` |
|
||||
| `MINT_OPERATION_TIMEOUT_SECONDS` | Per-attempt timeout for mint network calls | `30` |
|
||||
| `MINT_MAX_CONCURRENCY` | Concurrent operations allowed per mint (`0` disables the limit) | `4` |
|
||||
| `MINT_RETRY_MAX_ATTEMPTS` | Retries after a timeout or HTTP 429 (`0` disables retries) | `3` |
|
||||
| `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — |
|
||||
| `MIN_PAYOUT_SAT` | Min payout balance in sats (applies to all mints) | `210` |
|
||||
| `PAYOUT_INTERVAL_SECONDS` | Payout loop interval (seconds) | `900` |
|
||||
| `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` |
|
||||
| `CORS_ORIGINS` | Allowed CORS origins | `*` |
|
||||
| `RELAYS` | Nostr relays (comma-separated) | (default set) |
|
||||
| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to pause the refresh (previously discovered paths keep being served) | `600` |
|
||||
| `ENABLE_MODEL_PATHS_REFRESH` | Kill switch for the background model-path refresh (OpenRouter endpoint fan-out) | `true` |
|
||||
|
||||
Mint HTTP 429 responses create a per-mint cooldown. Operations that already hold
|
||||
Routstr's wallet mutation lock fail fast during that cooldown instead of waiting
|
||||
while blocking every other wallet mutation. Callers receive an error and may retry
|
||||
later; the current response does not include the cooldown duration.
|
||||
|
||||
### Priority
|
||||
|
||||
Environment variables are read on startup. Dashboard settings override them and persist in the database. Once you change a setting in the dashboard, the env var is ignored for that setting.
|
||||
|
||||
### 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
|
||||
@@ -156,3 +248,8 @@ Manage which AI models you offer:
|
||||
- **Create aliases** — friendly names for models
|
||||
|
||||
See [Pricing](pricing.md) for per-model pricing strategies.
|
||||
|
||||
Model path discovery is refreshed in the background and exposed through
|
||||
`/v1/models/paths`. The response groups each client-visible model ID with the
|
||||
provider paths that may appear in chat-completion response metadata. Tune the
|
||||
refresh cadence with `MODEL_PATHS_REFRESH_INTERVAL_SECONDS`.
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -395,8 +395,10 @@ and `routstr/upstream/ehbp.py`.
|
||||
`TINFOIL_API_KEY` env var.
|
||||
|
||||
- `routstr/upstream/ehbp.py`:
|
||||
- `parse_tinfoil_usage_metrics()` parses `prompt=N,completion=N[,total=N]`
|
||||
into an OpenAI-style usage dict.
|
||||
- `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`,
|
||||
@@ -404,11 +406,15 @@ and `routstr/upstream/ehbp.py`.
|
||||
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.
|
||||
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.
|
||||
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.
|
||||
@@ -448,13 +454,40 @@ 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.
|
||||
- Streaming requests: usage is delivered as an HTTP trailer. Currently the
|
||||
bearer path finalizes max-cost before streaming begins. Supporting streaming
|
||||
usage would require buffering the response (for X-Cashu) or a deferred
|
||||
finalization (for bearer).
|
||||
- ~~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.
|
||||
headers or trailers.
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
"""add model paths table
|
||||
|
||||
Revision ID: 64ed5594df1f
|
||||
Revises: aa50fde387a2
|
||||
Create Date: 2026-08-02 22:26:33.280409
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "64ed5594df1f"
|
||||
down_revision = "aa50fde387a2"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"model_paths",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||
sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||
sa.Column("provider_slug", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||
sa.Column("provider_type", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||
sa.Column("endpoint_tag", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||
sa.Column("endpoint_name", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||
sa.Column("upstream_provider_id", sa.Integer(), nullable=False),
|
||||
sa.Column("updated_at", sa.Integer(), nullable=False),
|
||||
sa.ForeignKeyConstraint(
|
||||
["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE"
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint(
|
||||
"model_id",
|
||||
"path",
|
||||
"upstream_provider_id",
|
||||
name="uq_model_paths_model_path_provider",
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_model_paths_upstream_provider_id"),
|
||||
"model_paths",
|
||||
["upstream_provider_id"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths")
|
||||
op.drop_table("model_paths")
|
||||
@@ -0,0 +1,62 @@
|
||||
"""add reservation release idempotency records
|
||||
|
||||
Revision ID: 7f2843d3f4e4
|
||||
Revises: fc4fa29630d2
|
||||
Create Date: 2026-07-24 02:06:06.066726
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "7f2843d3f4e4"
|
||||
down_revision = "fc4fa29630d2"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"reservation_releases",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("key_hash", sa.String(), nullable=False),
|
||||
sa.Column("billing_key_hash", sa.String(), nullable=False),
|
||||
sa.Column("reserved_msats", sa.Integer(), nullable=False),
|
||||
sa.Column(
|
||||
"status", sa.String(), nullable=False, server_default="active"
|
||||
),
|
||||
sa.Column("created_at", sa.Integer(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_reservation_releases_key_hash",
|
||||
"reservation_releases",
|
||||
["key_hash"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_reservation_releases_billing_key_hash",
|
||||
"reservation_releases",
|
||||
["billing_key_hash"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_reservation_releases_status_created_at",
|
||||
"reservation_releases",
|
||||
["status", "created_at"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index(
|
||||
"ix_reservation_releases_status_created_at",
|
||||
table_name="reservation_releases",
|
||||
)
|
||||
op.drop_index(
|
||||
"ix_reservation_releases_billing_key_hash",
|
||||
table_name="reservation_releases",
|
||||
)
|
||||
op.drop_index(
|
||||
"ix_reservation_releases_key_hash",
|
||||
table_name="reservation_releases",
|
||||
)
|
||||
op.drop_table("reservation_releases")
|
||||
@@ -0,0 +1,47 @@
|
||||
"""repair missing fee payout checkpoint columns
|
||||
|
||||
Revision ID: 9c4d8e2f1a6b
|
||||
Revises: 7f2843d3f4e4
|
||||
Create Date: 2026-07-25 00:00:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "9c4d8e2f1a6b"
|
||||
down_revision = "7f2843d3f4e4"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Repair databases stamped past the original checkpoint migration."""
|
||||
conn = op.get_bind()
|
||||
columns = {
|
||||
column["name"] for column in sa.inspect(conn).get_columns("routstr_fees")
|
||||
}
|
||||
|
||||
if "payout_in_progress_msats" not in columns:
|
||||
op.add_column(
|
||||
"routstr_fees",
|
||||
sa.Column(
|
||||
"payout_in_progress_msats",
|
||||
sa.Integer(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
),
|
||||
)
|
||||
|
||||
if "payout_started_at" not in columns:
|
||||
op.add_column(
|
||||
"routstr_fees",
|
||||
sa.Column("payout_started_at", sa.Integer(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# The preceding revision already expects both columns. This migration only
|
||||
# repairs schema drift, so downgrading it must preserve the expected schema.
|
||||
pass
|
||||
@@ -0,0 +1,39 @@
|
||||
"""add refund sweep claim lease
|
||||
|
||||
Revision ID: aa50fde387a2
|
||||
Revises: 9c4d8e2f1a6b
|
||||
Create Date: 2026-07-26 12:50:10.509217
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "aa50fde387a2"
|
||||
down_revision = "9c4d8e2f1a6b"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
columns = {
|
||||
column["name"]
|
||||
for column in sa.inspect(conn).get_columns("cashu_transactions")
|
||||
}
|
||||
if "sweep_started_at" not in columns:
|
||||
op.add_column(
|
||||
"cashu_transactions",
|
||||
sa.Column("sweep_started_at", sa.Integer(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
columns = {
|
||||
column["name"]
|
||||
for column in sa.inspect(conn).get_columns("cashu_transactions")
|
||||
}
|
||||
if "sweep_started_at" in columns:
|
||||
op.drop_column("cashu_transactions", "sweep_started_at")
|
||||
@@ -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,67 @@
|
||||
"""add mint url to lightning invoices
|
||||
|
||||
Revision ID: ecfa0d6e2a36
|
||||
Revises: 64ed5594df1f
|
||||
Create Date: 2026-08-02 23:53:00.037456
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "ecfa0d6e2a36"
|
||||
down_revision = "64ed5594df1f"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _resolve_backfill_mint_url(bind: sa.engine.Connection) -> str | None:
|
||||
"""Best-effort resolution of the mint that issued pre-existing invoices.
|
||||
|
||||
Order: persisted settings JSON -> PRIMARY_MINT_URL env -> first CASHU_MINTS entry.
|
||||
"""
|
||||
try:
|
||||
row = bind.execute(
|
||||
sa.text("SELECT data FROM settings ORDER BY id LIMIT 1")
|
||||
).fetchone()
|
||||
if row and row[0]:
|
||||
data = json.loads(row[0])
|
||||
mint = data.get("primary_mint") or next(
|
||||
iter(data.get("cashu_mints") or []), None
|
||||
)
|
||||
if mint:
|
||||
return str(mint)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
env_mint = os.environ.get("PRIMARY_MINT_URL", "").strip()
|
||||
if env_mint:
|
||||
return env_mint
|
||||
|
||||
cashu_mints = os.environ.get("CASHU_MINTS", "").strip()
|
||||
if cashu_mints:
|
||||
return cashu_mints.split(",")[0].strip() or None
|
||||
return None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"lightning_invoices", sa.Column("mint_url", sa.String(), nullable=True)
|
||||
)
|
||||
|
||||
bind = op.get_bind()
|
||||
backfill_mint = _resolve_backfill_mint_url(bind)
|
||||
if backfill_mint:
|
||||
bind.execute(
|
||||
sa.text(
|
||||
"UPDATE lightning_invoices SET mint_url = :mint WHERE mint_url IS NULL"
|
||||
),
|
||||
{"mint": backfill_mint},
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("lightning_invoices", "mint_url")
|
||||
@@ -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")
|
||||
+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
|
||||
|
||||
|
||||
+504
-175
@@ -3,16 +3,25 @@ import hashlib
|
||||
import math
|
||||
import random
|
||||
import time
|
||||
import uuid
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import case
|
||||
from sqlalchemy import case, inspect
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlmodel import col, select, update
|
||||
|
||||
from .core import get_logger
|
||||
from .core.db import ApiKey, AsyncSession, accumulate_routstr_fee
|
||||
from .core.db import (
|
||||
ApiKey,
|
||||
AsyncSession,
|
||||
ReservationRelease,
|
||||
accumulate_routstr_fee,
|
||||
create_session,
|
||||
)
|
||||
from .core.settings import settings
|
||||
from .payment.cost_calculation import (
|
||||
CostData,
|
||||
@@ -20,17 +29,65 @@ from .payment.cost_calculation import (
|
||||
MaxCostData,
|
||||
calculate_cost,
|
||||
)
|
||||
from .wallet import credit_balance, deserialize_token_from_string
|
||||
from .wallet import (
|
||||
classify_redemption_error,
|
||||
credit_balance,
|
||||
deserialize_token_from_string,
|
||||
wallet_operation_guard,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .payment.models import Model
|
||||
|
||||
logger = get_logger(__name__)
|
||||
payments_logger = get_logger("routstr.payments")
|
||||
|
||||
# Routstr platform fee constants
|
||||
ROUTSTR_FEE_PERCENT: float = 2.1
|
||||
ROUTSTR_LN_ADDRESS: str = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
|
||||
ROUTSTR_LN_ADDRESS: str = (
|
||||
"npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
|
||||
)
|
||||
ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900
|
||||
ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200
|
||||
|
||||
|
||||
def _format_msat_amount(amount: int) -> str:
|
||||
sats = f"{amount / 1000:.3f}".rstrip("0").rstrip(".")
|
||||
return f"{sats} sats ({amount} msats)"
|
||||
|
||||
|
||||
def _model_balance_error(required: int, available: int) -> dict[str, dict[str, str]]:
|
||||
return {
|
||||
"error": {
|
||||
"message": (
|
||||
f"Insufficient balance: {_format_msat_amount(required)} required "
|
||||
f"for this model; {_format_msat_amount(available)} available."
|
||||
),
|
||||
"type": "insufficient_quota",
|
||||
"code": "insufficient_balance",
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ReservationSnapshot:
|
||||
release_id: str
|
||||
key_hash: str
|
||||
billing_key_hash: str
|
||||
reserved_msats: int
|
||||
|
||||
|
||||
_current_reservation: ContextVar[ReservationSnapshot | None] = ContextVar(
|
||||
"current_billing_reservation", default=None
|
||||
)
|
||||
|
||||
|
||||
def _clear_current_reservation(snapshot: ReservationSnapshot) -> None:
|
||||
current = _current_reservation.get()
|
||||
if current is not None and current.release_id == snapshot.release_id:
|
||||
_current_reservation.set(None)
|
||||
|
||||
|
||||
# TODO: implement prepaid api key (not like it was before)
|
||||
# PREPAID_API_KEY = os.environ.get("PREPAID_API_KEY", None)
|
||||
# PREPAID_BALANCE = int(os.environ.get("PREPAID_BALANCE", "0")) * 1000 # Convert to msats
|
||||
@@ -78,12 +135,66 @@ async def check_and_reset_limit(key: ApiKey, session: AsyncSession) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def redemption_error_to_http_exception(error: Exception) -> HTTPException:
|
||||
"""Map a Cashu token redemption failure to a sanitized client-facing error.
|
||||
|
||||
Thin wrapper over the shared :func:`classify_redemption_error` so the bearer
|
||||
path stays identical to the X-Cashu and top-up paths.
|
||||
"""
|
||||
classified = classify_redemption_error(error)
|
||||
if classified is None:
|
||||
return HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": {
|
||||
"message": "Internal error during token redemption",
|
||||
"type": "api_error",
|
||||
"code": "internal_error",
|
||||
}
|
||||
},
|
||||
)
|
||||
error_type, status_code, message, error_code = classified
|
||||
return HTTPException(
|
||||
status_code=status_code,
|
||||
detail={
|
||||
"error": {
|
||||
"message": message,
|
||||
"type": error_type,
|
||||
"code": error_code,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def validate_bearer_key(
|
||||
bearer_key: str,
|
||||
session: AsyncSession,
|
||||
refund_address: Optional[str] = None,
|
||||
key_expiry_time: Optional[int] = None,
|
||||
min_cost: int = 0,
|
||||
) -> ApiKey:
|
||||
if bearer_key.startswith("cashu"):
|
||||
# Acquire before the first lookup/flush so concurrent token creation
|
||||
# cannot hold SQLite write transactions while waiting to mutate proofs.
|
||||
async with wallet_operation_guard():
|
||||
return await _validate_bearer_key_locked(
|
||||
bearer_key,
|
||||
session,
|
||||
refund_address,
|
||||
key_expiry_time,
|
||||
min_cost,
|
||||
)
|
||||
return await _validate_bearer_key_locked(
|
||||
bearer_key, session, refund_address, key_expiry_time, min_cost
|
||||
)
|
||||
|
||||
|
||||
async def _validate_bearer_key_locked(
|
||||
bearer_key: str,
|
||||
session: AsyncSession,
|
||||
refund_address: Optional[str] = None,
|
||||
key_expiry_time: Optional[int] = None,
|
||||
min_cost: int = 0,
|
||||
) -> ApiKey:
|
||||
"""
|
||||
Validates the provided API key using SQLModel.
|
||||
@@ -171,13 +282,7 @@ async def validate_bearer_key(
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Insufficient balance: {min_cost} mSats required for this model. {billing_key.total_balance} available.",
|
||||
"type": "insufficient_quota",
|
||||
"code": "insufficient_balance",
|
||||
}
|
||||
},
|
||||
detail=_model_balance_error(min_cost, billing_key.total_balance),
|
||||
)
|
||||
|
||||
# Early check: Spending limit check (Child key limit)
|
||||
@@ -216,7 +321,17 @@ async def validate_bearer_key(
|
||||
|
||||
try:
|
||||
hashed_key = hashlib.sha256(bearer_key.encode()).hexdigest()
|
||||
token_obj = deserialize_token_from_string(bearer_key)
|
||||
try:
|
||||
token_obj = deserialize_token_from_string(bearer_key)
|
||||
except Exception as decode_error:
|
||||
# A malformed token is a bad token (400 invalid_cashu_token via
|
||||
# the shared taxonomy), not an auth failure (401) — otherwise it
|
||||
# would fall through to the generic "Invalid API key" handler.
|
||||
raise redemption_error_to_http_exception(
|
||||
ValueError(
|
||||
f"Invalid Cashu token: could not decode token ({decode_error})"
|
||||
)
|
||||
) from decode_error
|
||||
logger.debug(
|
||||
"Generated token hash", extra={"hash_preview": hashed_key[:16] + "..."}
|
||||
)
|
||||
@@ -257,13 +372,9 @@ async def validate_bearer_key(
|
||||
if min_cost > 0 and existing_key.total_balance < min_cost:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Insufficient balance: {min_cost} mSats required for this model. {existing_key.total_balance} available.",
|
||||
"type": "insufficient_quota",
|
||||
"code": "insufficient_balance",
|
||||
}
|
||||
},
|
||||
detail=_model_balance_error(
|
||||
min_cost, existing_key.total_balance
|
||||
),
|
||||
)
|
||||
|
||||
return existing_key
|
||||
@@ -276,11 +387,23 @@ async def validate_bearer_key(
|
||||
"has_expiry_time": bool(key_expiry_time),
|
||||
},
|
||||
)
|
||||
if token_obj.mint in settings.cashu_mints:
|
||||
if token_obj.mint == settings.primary_mint:
|
||||
if token_obj.unit != settings.primary_mint_unit:
|
||||
raise redemption_error_to_http_exception(
|
||||
ValueError(
|
||||
"Cashu token unit does not match the configured primary "
|
||||
f"mint unit: expected {settings.primary_mint_unit}, "
|
||||
f"got {token_obj.unit}"
|
||||
)
|
||||
)
|
||||
refund_currency = token_obj.unit
|
||||
refund_mint_url = settings.primary_mint
|
||||
elif token_obj.mint in settings.cashu_mints:
|
||||
refund_currency = token_obj.unit
|
||||
refund_mint_url = token_obj.mint
|
||||
else:
|
||||
refund_currency = "sat"
|
||||
# Foreign tokens are swapped into the configured primary mint.
|
||||
refund_currency = settings.primary_mint_unit
|
||||
refund_mint_url = settings.primary_mint
|
||||
|
||||
new_key = ApiKey(
|
||||
@@ -334,19 +457,32 @@ async def validate_bearer_key(
|
||||
"error_type": type(credit_error).__name__,
|
||||
},
|
||||
)
|
||||
raise credit_error
|
||||
await session.rollback()
|
||||
raise redemption_error_to_http_exception(credit_error) from credit_error
|
||||
|
||||
if msats <= 0:
|
||||
logger.error(
|
||||
"Token redemption returned zero or negative amount",
|
||||
extra={"msats": msats, "key_hash": hashed_key[:8] + "..."},
|
||||
)
|
||||
# Defense-in-depth: credit_balance now refuses to commit on a
|
||||
# zero/negative redemption, but if a row was nonetheless
|
||||
# persisted, drop it so we never leave an orphan zero-balance key.
|
||||
# Defense-in-depth: credit_balance already raises
|
||||
# ValueError("Redeemed token amount must be positive…") before
|
||||
# returning (wallet.py), so this branch is only reachable if a
|
||||
# zero/negative row was somehow persisted; drop it so we never
|
||||
# leave an orphan zero-balance key. Reuse the shared taxonomy
|
||||
# (cashu_error) so the envelope matches the mapper above.
|
||||
await session.delete(new_key)
|
||||
await session.commit()
|
||||
raise Exception("Token redemption failed")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": {
|
||||
"message": "Failed to redeem Cashu token: token yielded no value",
|
||||
"type": "cashu_error",
|
||||
"code": "cashu_token_zero_value",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
await session.refresh(new_key)
|
||||
await session.commit()
|
||||
@@ -379,7 +515,7 @@ async def validate_bearer_key(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Invalid or expired Cashu key: {str(e)}",
|
||||
"message": "Invalid or expired Cashu key",
|
||||
"type": "invalid_request_error",
|
||||
"code": "invalid_api_key",
|
||||
}
|
||||
@@ -534,6 +670,16 @@ async def pay_for_request(
|
||||
},
|
||||
)
|
||||
|
||||
# Create the durable reservation identity before changing aggregate balances.
|
||||
# The row and balance updates commit together, so every reserved amount has one
|
||||
# owner that can reach exactly one terminal state.
|
||||
reservation = ReservationSnapshot(
|
||||
release_id=uuid.uuid4().hex,
|
||||
key_hash=key.hashed_key,
|
||||
billing_key_hash=billing_key.hashed_key,
|
||||
reserved_msats=cost_per_request,
|
||||
)
|
||||
|
||||
# Charge the base cost for the request atomically to avoid race conditions
|
||||
reserved_at_now = int(time.time())
|
||||
stmt = (
|
||||
@@ -595,22 +741,83 @@ async def pay_for_request(
|
||||
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
|
||||
|
||||
if child_result.rowcount == 0:
|
||||
# Build the error before rollback expires ORM attributes.
|
||||
limit_message = (
|
||||
f"Balance limit exceeded: {key.balance_limit} mSats limit. "
|
||||
f"{key.total_spent} already spent ({key.reserved_balance} reserved), "
|
||||
f"{cost_per_request} required for this request."
|
||||
)
|
||||
# The parent reservation update already ran in this transaction.
|
||||
# Roll it back before failover code attempts to restore the previous
|
||||
# reservation; otherwise that later commit can persist both updates.
|
||||
await session.rollback()
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Balance limit exceeded: {key.balance_limit} mSats limit. {key.total_spent} already spent ({key.reserved_balance} reserved), {cost_per_request} required for this request.",
|
||||
"message": limit_message,
|
||||
"type": "insufficient_quota",
|
||||
"code": "balance_limit_exceeded",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
await session.commit()
|
||||
session.add(
|
||||
ReservationRelease(
|
||||
id=reservation.release_id,
|
||||
key_hash=reservation.key_hash,
|
||||
billing_key_hash=reservation.billing_key_hash,
|
||||
reserved_msats=reservation.reserved_msats,
|
||||
status="active",
|
||||
)
|
||||
)
|
||||
# Publish the identity before commit. If the commit succeeds but its
|
||||
# acknowledgement is interrupted, exact cleanup can still recover the
|
||||
# durable row. A definitely failed commit is harmless because every
|
||||
# terminal transition validates that row before touching balances.
|
||||
_current_reservation.set(reservation)
|
||||
try:
|
||||
await session.commit()
|
||||
except BaseException:
|
||||
# The database may have committed even if acknowledgement was cancelled
|
||||
# or the connection failed. Reconcile using a fresh transaction and the
|
||||
# exact durable identity; no upstream request has started yet.
|
||||
try:
|
||||
await session.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
async with create_session() as cleanup_session:
|
||||
record = await cleanup_session.get(
|
||||
ReservationRelease, reservation.release_id
|
||||
)
|
||||
if record is not None and record.status == "active":
|
||||
await _transition_reservation_to_released(
|
||||
reservation,
|
||||
cleanup_session,
|
||||
decrement_requests=True,
|
||||
idempotent_success=True,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to reconcile ambiguous reservation commit",
|
||||
extra={"reservation_id": reservation.release_id},
|
||||
)
|
||||
finally:
|
||||
_clear_current_reservation(reservation)
|
||||
raise
|
||||
|
||||
await session.refresh(billing_key)
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
await session.refresh(key)
|
||||
try:
|
||||
await session.refresh(billing_key)
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
await session.refresh(key)
|
||||
except Exception:
|
||||
# The reservation transaction is already committed and durable. Logging
|
||||
# refresh failures must not make the caller treat it as unreserved.
|
||||
logger.exception(
|
||||
"Reservation committed but post-commit refresh failed",
|
||||
extra={"reservation_id": reservation.release_id},
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Payment processed successfully",
|
||||
@@ -640,89 +847,206 @@ async def pay_for_request(
|
||||
|
||||
|
||||
async def revert_pay_for_request(
|
||||
key: ApiKey, session: AsyncSession, cost_per_request: int
|
||||
key: ApiKey,
|
||||
session: AsyncSession,
|
||||
cost_per_request: int,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
) -> bool:
|
||||
"""Revert a previously reserved payment. Returns True if revert succeeded,
|
||||
False if the reservation was already released (prevents negative reserved_balance)."""
|
||||
billing_key = await get_billing_key(key, session)
|
||||
|
||||
# Keep reserved_at while other reservations remain
|
||||
cleared_reserved_at = case(
|
||||
(col(ApiKey.reserved_balance) - cost_per_request > 0, col(ApiKey.reserved_at)),
|
||||
else_=None,
|
||||
"""Revert the current request's durable reservation exactly once."""
|
||||
snapshot = reservation_snapshot or await get_reservation_snapshot(key, session)
|
||||
await _validate_reservation_snapshot(key, snapshot, session, require_active=False)
|
||||
if cost_per_request != snapshot.reserved_msats:
|
||||
return False
|
||||
return await _transition_reservation_to_released(
|
||||
snapshot,
|
||||
session,
|
||||
decrement_requests=True,
|
||||
idempotent_success=False,
|
||||
)
|
||||
|
||||
stmt = (
|
||||
|
||||
async def _validate_reservation_snapshot(
|
||||
key: ApiKey,
|
||||
snapshot: ReservationSnapshot,
|
||||
session: AsyncSession,
|
||||
*,
|
||||
require_active: bool = True,
|
||||
) -> None:
|
||||
"""Reject cross-request or forged reservation handles before any mutation."""
|
||||
state = inspect(key)
|
||||
identity = state.identity if state is not None else None
|
||||
key_hash = str(identity[0]) if identity else key.__dict__.get("hashed_key")
|
||||
if snapshot.key_hash != key_hash:
|
||||
raise RuntimeError("Billing reservation does not belong to this key")
|
||||
|
||||
persisted_key = await session.get(ApiKey, snapshot.key_hash)
|
||||
if persisted_key is None:
|
||||
raise RuntimeError("Billing reservation key no longer exists")
|
||||
expected_billing_hash = persisted_key.parent_key_hash or persisted_key.hashed_key
|
||||
if snapshot.billing_key_hash != expected_billing_hash:
|
||||
raise RuntimeError("Billing reservation does not belong to this billing key")
|
||||
|
||||
record = await session.get(ReservationRelease, snapshot.release_id)
|
||||
if (
|
||||
record is None
|
||||
or (require_active and record.status != "active")
|
||||
or record.key_hash != snapshot.key_hash
|
||||
or record.billing_key_hash != snapshot.billing_key_hash
|
||||
or record.reserved_msats != snapshot.reserved_msats
|
||||
):
|
||||
raise RuntimeError("Billing reservation record does not match the request")
|
||||
|
||||
|
||||
async def get_reservation_snapshot(
|
||||
key: ApiKey, session: AsyncSession
|
||||
) -> ReservationSnapshot:
|
||||
"""Return the durable reservation created for the current request."""
|
||||
snapshot = _current_reservation.get()
|
||||
if snapshot is None:
|
||||
raise RuntimeError("No billing reservation is associated with this request")
|
||||
await _validate_reservation_snapshot(key, snapshot, session)
|
||||
return snapshot
|
||||
|
||||
|
||||
async def _transition_reservation_to_released(
|
||||
snapshot: ReservationSnapshot,
|
||||
session: AsyncSession,
|
||||
*,
|
||||
decrement_requests: bool,
|
||||
idempotent_success: bool,
|
||||
) -> bool:
|
||||
transition = (
|
||||
update(ReservationRelease)
|
||||
.where(col(ReservationRelease.id) == snapshot.release_id)
|
||||
.where(col(ReservationRelease.status) == "active")
|
||||
.where(col(ReservationRelease.key_hash) == snapshot.key_hash)
|
||||
.where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash)
|
||||
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
|
||||
.values(status="released")
|
||||
)
|
||||
transition_result = await session.exec(transition) # type: ignore[call-overload]
|
||||
if transition_result.rowcount != 1:
|
||||
await session.rollback()
|
||||
existing = await session.get(ReservationRelease, snapshot.release_id)
|
||||
return bool(
|
||||
idempotent_success
|
||||
and existing is not None
|
||||
and existing.status == "released"
|
||||
and existing.key_hash == snapshot.key_hash
|
||||
and existing.billing_key_hash == snapshot.billing_key_hash
|
||||
and existing.reserved_msats == snapshot.reserved_msats
|
||||
)
|
||||
|
||||
values: dict[str, object] = {
|
||||
"reserved_balance": col(ApiKey.reserved_balance) - snapshot.reserved_msats,
|
||||
"reserved_at": case(
|
||||
(
|
||||
col(ApiKey.reserved_balance) - snapshot.reserved_msats > 0,
|
||||
col(ApiKey.reserved_at),
|
||||
),
|
||||
else_=None,
|
||||
),
|
||||
}
|
||||
if decrement_requests:
|
||||
values["total_requests"] = col(ApiKey.total_requests) - 1
|
||||
|
||||
release_stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
||||
.where(col(ApiKey.reserved_balance) >= cost_per_request)
|
||||
.values(
|
||||
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
|
||||
reserved_at=cleared_reserved_at,
|
||||
total_requests=col(ApiKey.total_requests) - 1,
|
||||
)
|
||||
.where(col(ApiKey.hashed_key) == snapshot.billing_key_hash)
|
||||
.where(col(ApiKey.reserved_balance) >= snapshot.reserved_msats)
|
||||
.values(**values)
|
||||
)
|
||||
result = await session.exec(release_stmt) # type: ignore[call-overload]
|
||||
if result.rowcount != 1:
|
||||
await session.rollback()
|
||||
return False
|
||||
|
||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
||||
|
||||
# Also decrement total_requests and reserved_balance on the child key if it's different
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
child_stmt = (
|
||||
if snapshot.billing_key_hash != snapshot.key_hash:
|
||||
child_release_stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
||||
.where(col(ApiKey.reserved_balance) >= cost_per_request)
|
||||
.values(
|
||||
total_requests=col(ApiKey.total_requests) - 1,
|
||||
reserved_balance=col(ApiKey.reserved_balance) - cost_per_request,
|
||||
reserved_at=cleared_reserved_at,
|
||||
)
|
||||
.where(col(ApiKey.hashed_key) == snapshot.key_hash)
|
||||
.where(col(ApiKey.reserved_balance) >= snapshot.reserved_msats)
|
||||
.values(**values)
|
||||
)
|
||||
await session.exec(child_stmt) # type: ignore[call-overload]
|
||||
child_result = await session.exec( # type: ignore[call-overload]
|
||||
child_release_stmt
|
||||
)
|
||||
if child_result.rowcount != 1:
|
||||
await session.rollback()
|
||||
return False
|
||||
|
||||
await session.commit()
|
||||
if result.rowcount == 0:
|
||||
logger.warning(
|
||||
"Revert skipped - reservation already released (no-op to prevent negative reserved_balance)",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
||||
"cost_to_revert": cost_per_request,
|
||||
"current_reserved_balance": billing_key.reserved_balance,
|
||||
},
|
||||
)
|
||||
return False
|
||||
await session.refresh(billing_key)
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
await session.refresh(key)
|
||||
payments_logger.info(
|
||||
"REVERT",
|
||||
extra={
|
||||
"event": "revert",
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
||||
"cost_reverted": cost_per_request,
|
||||
"balance": billing_key.balance,
|
||||
"reserved_balance": billing_key.reserved_balance,
|
||||
},
|
||||
)
|
||||
_clear_current_reservation(snapshot)
|
||||
return True
|
||||
|
||||
|
||||
async def release_reservation(
|
||||
snapshot: ReservationSnapshot,
|
||||
session: AsyncSession,
|
||||
reserved_msats: int,
|
||||
) -> bool:
|
||||
"""Release one durable reservation exactly once without charging."""
|
||||
if reserved_msats <= 0 or reserved_msats != snapshot.reserved_msats:
|
||||
return False
|
||||
return await _transition_reservation_to_released(
|
||||
snapshot,
|
||||
session,
|
||||
decrement_requests=False,
|
||||
idempotent_success=True,
|
||||
)
|
||||
|
||||
|
||||
async def _claim_reservation_for_charge(
|
||||
snapshot: ReservationSnapshot, session: AsyncSession
|
||||
) -> bool:
|
||||
"""Claim an active reservation in the caller's charge transaction."""
|
||||
statement = (
|
||||
update(ReservationRelease)
|
||||
.where(col(ReservationRelease.id) == snapshot.release_id)
|
||||
.where(col(ReservationRelease.status) == "active")
|
||||
.where(col(ReservationRelease.key_hash) == snapshot.key_hash)
|
||||
.where(col(ReservationRelease.billing_key_hash) == snapshot.billing_key_hash)
|
||||
.where(col(ReservationRelease.reserved_msats) == snapshot.reserved_msats)
|
||||
.values(status="charged")
|
||||
)
|
||||
result = await session.exec(statement) # type: ignore[call-overload]
|
||||
if result.rowcount == 1:
|
||||
_clear_current_reservation(snapshot)
|
||||
return True
|
||||
|
||||
await session.rollback()
|
||||
return False
|
||||
|
||||
|
||||
async def adjust_payment_for_tokens(
|
||||
key: ApiKey,
|
||||
response_data: dict,
|
||||
session: AsyncSession,
|
||||
deducted_max_cost: int,
|
||||
model_obj: "Model | None" = None,
|
||||
provider_fee: float | None = None,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
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``.
|
||||
"""
|
||||
billing_key = await get_billing_key(key, session)
|
||||
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
||||
await _validate_reservation_snapshot(
|
||||
key, reservation, session, require_active=False
|
||||
)
|
||||
# The persisted amount is authoritative if request-level minimum pricing
|
||||
# changed the caller's original estimate.
|
||||
deducted_max_cost = reservation.reserved_msats
|
||||
model = response_data.get("model", "unknown")
|
||||
|
||||
logger.debug(
|
||||
@@ -738,50 +1062,21 @@ async def adjust_payment_for_tokens(
|
||||
)
|
||||
|
||||
async def release_reservation_only() -> None:
|
||||
"""Fallback to release reservation without charging when main update fails."""
|
||||
"""Fallback to release this request's reservation without charging."""
|
||||
try:
|
||||
release_stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
||||
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
|
||||
.values(
|
||||
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost
|
||||
)
|
||||
released = await release_reservation(
|
||||
reservation, session, reservation.reserved_msats
|
||||
)
|
||||
logger.warning(
|
||||
"Released reservation without charging (fallback)"
|
||||
if released
|
||||
else "Reservation was already finalized; fallback skipped",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
||||
"deducted_max_cost": deducted_max_cost,
|
||||
},
|
||||
)
|
||||
result = await session.exec(release_stmt) # type: ignore[call-overload]
|
||||
|
||||
# Also release on child key if it's different
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
child_release_stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
||||
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
|
||||
.values(
|
||||
reserved_balance=col(ApiKey.reserved_balance)
|
||||
- deducted_max_cost
|
||||
)
|
||||
)
|
||||
await session.exec(child_release_stmt) # type: ignore[call-overload]
|
||||
|
||||
await session.commit()
|
||||
if result.rowcount == 0: # type: ignore[union-attr]
|
||||
logger.warning(
|
||||
"Release reservation skipped - already released (no-op to prevent negative reserved_balance)",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
||||
"deducted_max_cost": deducted_max_cost,
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Released reservation without charging (fallback)",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
||||
"deducted_max_cost": deducted_max_cost,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to release reservation in fallback",
|
||||
@@ -803,7 +1098,17 @@ async def adjust_payment_for_tokens(
|
||||
extra={"error": str(e), "fee_msats": fee_msats},
|
||||
)
|
||||
|
||||
match await calculate_cost(response_data, deducted_max_cost):
|
||||
calculated_cost = await calculate_cost(
|
||||
response_data, deducted_max_cost, model_obj, provider_fee
|
||||
)
|
||||
if not isinstance(calculated_cost, CostDataError):
|
||||
if not await _claim_reservation_for_charge(reservation, session):
|
||||
# A prior charge or release already owns this reservation. Returning
|
||||
# the calculated metadata is safe; the aggregate balances must not
|
||||
# be modified a second time.
|
||||
return calculated_cost.dict()
|
||||
|
||||
match calculated_cost:
|
||||
case MaxCostData() as cost:
|
||||
logger.debug(
|
||||
"Using max cost data (no token adjustment)",
|
||||
@@ -831,8 +1136,10 @@ async def adjust_payment_for_tokens(
|
||||
)
|
||||
|
||||
safe_reserved = case(
|
||||
(col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost),
|
||||
(
|
||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
|
||||
@@ -850,8 +1157,10 @@ async def adjust_payment_for_tokens(
|
||||
# Also update total_spent and reserved_balance on the child key if it's different
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
child_safe_reserved = case(
|
||||
(col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost),
|
||||
(
|
||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
child_stmt = (
|
||||
@@ -961,8 +1270,10 @@ async def adjust_payment_for_tokens(
|
||||
)
|
||||
|
||||
exact_safe_reserved = case(
|
||||
(col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost),
|
||||
(
|
||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
|
||||
@@ -980,8 +1291,10 @@ async def adjust_payment_for_tokens(
|
||||
# Also update total_spent and reserved_balance on the child key if it's different
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
child_exact_safe_reserved = case(
|
||||
(col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost),
|
||||
(
|
||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
child_stmt = (
|
||||
@@ -1020,31 +1333,45 @@ async def adjust_payment_for_tokens(
|
||||
|
||||
# actual cost exceeded discounted reservation (due to tolerance_percentage)
|
||||
if cost_difference > 0:
|
||||
# Always release the reservation and charge min(actual_cost, balance).
|
||||
# CASE expressions keep this atomic and safe even when the
|
||||
# stale-reservation sweeper has already released the reservation.
|
||||
chargeable = case(
|
||||
(col(ApiKey.balance) >= total_cost_msats, total_cost_msats),
|
||||
else_=col(ApiKey.balance),
|
||||
)
|
||||
overrun_safe_reserved = case(
|
||||
(
|
||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
|
||||
finalize_stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
||||
.values(
|
||||
reserved_balance=overrun_safe_reserved,
|
||||
balance=col(ApiKey.balance) - chargeable,
|
||||
total_spent=col(ApiKey.total_spent) + chargeable,
|
||||
# Lock the billing row so the parent and child record the same
|
||||
# database-determined charge under concurrent finalizations.
|
||||
actual_charge_msats = 0
|
||||
for attempt in range(5):
|
||||
locked_billing_key = (
|
||||
await session.exec(
|
||||
select(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
||||
.with_for_update()
|
||||
.execution_options(populate_existing=True)
|
||||
)
|
||||
).one()
|
||||
observed_balance = locked_billing_key.balance
|
||||
actual_charge_msats = min(observed_balance, total_cost_msats)
|
||||
overrun_safe_reserved = case(
|
||||
(
|
||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
)
|
||||
await session.exec(finalize_stmt) # type: ignore[call-overload]
|
||||
finalize_result = await session.exec( # type: ignore[call-overload]
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
||||
.where(col(ApiKey.balance) == observed_balance)
|
||||
.values(
|
||||
reserved_balance=overrun_safe_reserved,
|
||||
balance=col(ApiKey.balance) - actual_charge_msats,
|
||||
total_spent=col(ApiKey.total_spent) + actual_charge_msats,
|
||||
)
|
||||
)
|
||||
if finalize_result.rowcount == 1:
|
||||
break
|
||||
await session.rollback()
|
||||
if not await _claim_reservation_for_charge(reservation, session):
|
||||
return cost.dict()
|
||||
else:
|
||||
await session.rollback()
|
||||
raise RuntimeError("Could not atomically finalize cost overrun")
|
||||
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
child_stmt = (
|
||||
@@ -1052,7 +1379,7 @@ async def adjust_payment_for_tokens(
|
||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
||||
.values(
|
||||
reserved_balance=overrun_safe_reserved,
|
||||
total_spent=col(ApiKey.total_spent) + min(billing_key.balance, total_cost_msats),
|
||||
total_spent=col(ApiKey.total_spent) + actual_charge_msats,
|
||||
)
|
||||
)
|
||||
await session.exec(child_stmt) # type: ignore[call-overload]
|
||||
@@ -1062,18 +1389,18 @@ async def adjust_payment_for_tokens(
|
||||
await session.refresh(billing_key)
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
await session.refresh(key)
|
||||
cost.total_msats = total_cost_msats
|
||||
cost.total_msats = actual_charge_msats
|
||||
logger.info(
|
||||
"Finalized payment with additional charge",
|
||||
extra={
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
||||
"charged_amount": total_cost_msats,
|
||||
"charged_amount": actual_charge_msats,
|
||||
"new_balance": billing_key.balance,
|
||||
"model": model,
|
||||
},
|
||||
)
|
||||
await _accumulate_fee(total_cost_msats)
|
||||
await _accumulate_fee(actual_charge_msats)
|
||||
payments_logger.info(
|
||||
"FINALIZE",
|
||||
extra={
|
||||
@@ -1082,7 +1409,7 @@ async def adjust_payment_for_tokens(
|
||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
||||
"model": model,
|
||||
"cost_reserved": deducted_max_cost,
|
||||
"cost_charged": total_cost_msats,
|
||||
"cost_charged": actual_charge_msats,
|
||||
"input_tokens": cost.input_tokens,
|
||||
"output_tokens": cost.output_tokens,
|
||||
"balance": billing_key.balance,
|
||||
@@ -1122,8 +1449,10 @@ async def adjust_payment_for_tokens(
|
||||
)
|
||||
|
||||
refund_safe_reserved = case(
|
||||
(col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost),
|
||||
(
|
||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
|
||||
@@ -1141,8 +1470,10 @@ async def adjust_payment_for_tokens(
|
||||
# Also update total_spent and reserved_balance on the child key if it's different
|
||||
if billing_key.hashed_key != key.hashed_key:
|
||||
child_refund_safe_reserved = case(
|
||||
(col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost),
|
||||
(
|
||||
col(ApiKey.reserved_balance) >= deducted_max_cost,
|
||||
col(ApiKey.reserved_balance) - deducted_max_cost,
|
||||
),
|
||||
else_=0,
|
||||
)
|
||||
child_stmt = (
|
||||
@@ -1317,9 +1648,7 @@ async def periodic_dead_key_prune() -> None:
|
||||
|
||||
try:
|
||||
async with create_session() as session:
|
||||
await prune_dead_api_keys(
|
||||
session, settings.dead_key_min_age_seconds
|
||||
)
|
||||
await prune_dead_api_keys(session, settings.dead_key_min_age_seconds)
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
|
||||
+227
-68
@@ -7,7 +7,7 @@ from typing import Annotated, NoReturn
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import BaseModel
|
||||
from sqlmodel import col, or_, select, update
|
||||
from sqlmodel import col, select, update
|
||||
|
||||
from .auth import get_billing_key, validate_bearer_key
|
||||
from .core.db import (
|
||||
@@ -15,12 +15,24 @@ from .core.db import (
|
||||
AsyncSession,
|
||||
CashuTransaction,
|
||||
get_session,
|
||||
store_cashu_transaction,
|
||||
release_stale_reservations,
|
||||
)
|
||||
from .core.db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from .core.logging import get_logger
|
||||
from .core.settings import settings
|
||||
from .lightning import lightning_router
|
||||
from .wallet import credit_balance, recieve_token, send_to_lnurl, send_token
|
||||
from .payment.lnurl import MeltOutcomeAmbiguousError
|
||||
from .wallet import (
|
||||
classify_redemption_error,
|
||||
credit_balance,
|
||||
is_mint_connection_error,
|
||||
recieve_token,
|
||||
send_to_lnurl,
|
||||
send_token,
|
||||
token_mint_url,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
balance_router = APIRouter(prefix="/v1/balance")
|
||||
@@ -99,13 +111,19 @@ async def account_info(
|
||||
# Note: validate_bearer_key already supports refund_address and key_expiry_time params
|
||||
|
||||
|
||||
@router.get("/create")
|
||||
async def create_balance(
|
||||
class BalanceCreateRequest(BaseModel):
|
||||
initial_balance_token: str
|
||||
balance_limit: int | None = None
|
||||
balance_limit_reset: str | None = None
|
||||
validity_date: int | None = None
|
||||
|
||||
|
||||
async def _create_balance(
|
||||
initial_balance_token: str,
|
||||
balance_limit: int | None = None,
|
||||
balance_limit_reset: str | None = None,
|
||||
validity_date: int | None = None,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
balance_limit: int | None,
|
||||
balance_limit_reset: str | None,
|
||||
validity_date: int | None,
|
||||
session: AsyncSession,
|
||||
) -> dict:
|
||||
key = await validate_bearer_key(initial_balance_token, session)
|
||||
|
||||
@@ -125,6 +143,37 @@ async def create_balance(
|
||||
}
|
||||
|
||||
|
||||
@router.post("/create")
|
||||
async def create_balance_from_body(
|
||||
payload: BalanceCreateRequest,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
return await _create_balance(
|
||||
payload.initial_balance_token,
|
||||
payload.balance_limit,
|
||||
payload.balance_limit_reset,
|
||||
payload.validity_date,
|
||||
session,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/create")
|
||||
async def create_balance(
|
||||
initial_balance_token: str,
|
||||
balance_limit: int | None = None,
|
||||
balance_limit_reset: str | None = None,
|
||||
validity_date: int | None = None,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
return await _create_balance(
|
||||
initial_balance_token,
|
||||
balance_limit,
|
||||
balance_limit_reset,
|
||||
validity_date,
|
||||
session,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/info")
|
||||
async def wallet_info(
|
||||
key: ApiKey = Depends(get_key_from_header),
|
||||
@@ -137,6 +186,17 @@ class TopupRequest(BaseModel):
|
||||
cashu_token: str
|
||||
|
||||
|
||||
def _error_chain(error: BaseException) -> list[dict[str, str]]:
|
||||
chain: list[dict[str, str]] = []
|
||||
current: BaseException | None = error
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
chain.append({"type": type(current).__name__, "message": str(current)})
|
||||
current = current.__cause__ or current.__context__
|
||||
return chain
|
||||
|
||||
|
||||
@router.post("/topup")
|
||||
async def topup_wallet_endpoint(
|
||||
cashu_token: str | None = None,
|
||||
@@ -154,32 +214,61 @@ async def topup_wallet_endpoint(
|
||||
cashu_token = cashu_token.replace("\n", "").replace("\r", "").replace("\t", "")
|
||||
if len(cashu_token) < 10 or "cashu" not in cashu_token:
|
||||
raise HTTPException(status_code=400, detail="Invalid token format")
|
||||
|
||||
source_mint = token_mint_url(cashu_token, "unknown")
|
||||
logger.info(
|
||||
"Cashu wallet top-up started",
|
||||
extra={
|
||||
"event": "cashu_topup_started",
|
||||
"source_mint": source_mint,
|
||||
"primary_mint": settings.primary_mint,
|
||||
"trusted_mints": settings.cashu_mints,
|
||||
"key_hash": billing_key.hashed_key[:8],
|
||||
},
|
||||
)
|
||||
try:
|
||||
amount_msats = await credit_balance(cashu_token, billing_key, session)
|
||||
except ValueError as e:
|
||||
error_msg = str(e)
|
||||
if "already spent" in error_msg.lower():
|
||||
raise HTTPException(status_code=400, detail="Token already spent")
|
||||
elif "invalid" in error_msg.lower() or "decode" in error_msg.lower():
|
||||
raise HTTPException(status_code=400, detail="Invalid token format")
|
||||
elif "insufficient" in error_msg.lower() or "melt fee" in error_msg.lower():
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Token value is too small to cover swap fees. {error_msg}",
|
||||
)
|
||||
elif "failed to melt" in error_msg.lower():
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Failed to swap foreign mint token. {error_msg}",
|
||||
)
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail=f"Failed to redeem token: {error_msg}")
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"topup_wallet_endpoint: unhandled error",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
# Shared taxonomy so top-up matches the bearer/X-Cashu paths (503 for an
|
||||
# unreachable mint, 422 for fee/swap failures, 400 for token faults).
|
||||
classified = classify_redemption_error(e)
|
||||
if classified is None:
|
||||
logger.error(
|
||||
"Cashu wallet top-up failed with an unhandled error",
|
||||
extra={
|
||||
"event": "cashu_topup_failed",
|
||||
"source_mint": source_mint,
|
||||
"primary_mint": settings.primary_mint,
|
||||
"trusted_mints": settings.cashu_mints,
|
||||
"error_chain": _error_chain(e),
|
||||
},
|
||||
)
|
||||
raise HTTPException(status_code=500, detail="Internal server error")
|
||||
error_type, status_code, message, error_code = classified
|
||||
logger.warning(
|
||||
"Cashu wallet top-up failed",
|
||||
extra={
|
||||
"event": "cashu_topup_failed",
|
||||
"source_mint": source_mint,
|
||||
"primary_mint": settings.primary_mint,
|
||||
"trusted_mints": settings.cashu_mints,
|
||||
"status_code": status_code,
|
||||
"error_type": error_type,
|
||||
"error_code": error_code,
|
||||
"error_chain": _error_chain(e),
|
||||
},
|
||||
)
|
||||
raise HTTPException(status_code=500, detail="Internal server error")
|
||||
raise HTTPException(status_code=status_code, detail=message)
|
||||
|
||||
logger.info(
|
||||
"Cashu wallet top-up completed",
|
||||
extra={
|
||||
"event": "cashu_topup_completed",
|
||||
"source_mint": source_mint,
|
||||
"credited_msats": amount_msats,
|
||||
"key_hash": billing_key.hashed_key[:8],
|
||||
},
|
||||
)
|
||||
return {"msats": amount_msats}
|
||||
|
||||
|
||||
@@ -224,8 +313,42 @@ async def _lookup_key_no_create(
|
||||
return None
|
||||
|
||||
|
||||
async def _get_persisted_api_key_refund(
|
||||
key: ApiKey, session: AsyncSession
|
||||
) -> dict[str, str] | None:
|
||||
result = await session.exec(
|
||||
select(CashuTransaction)
|
||||
.where(
|
||||
CashuTransaction.api_key_hashed_key == key.hashed_key,
|
||||
CashuTransaction.type == "out",
|
||||
CashuTransaction.source == "apikey",
|
||||
)
|
||||
.order_by(col(CashuTransaction.created_at).desc())
|
||||
)
|
||||
refund = result.first()
|
||||
if refund is None:
|
||||
return None
|
||||
if refund.swept:
|
||||
raise HTTPException(status_code=410, detail="Refund has been swept")
|
||||
|
||||
refund.collected = True
|
||||
session.add(refund)
|
||||
await session.commit()
|
||||
|
||||
persisted = {"token": refund.token}
|
||||
if refund.unit == "sat":
|
||||
persisted["sats"] = str(refund.amount)
|
||||
else:
|
||||
persisted["msats"] = str(refund.amount)
|
||||
return persisted
|
||||
|
||||
|
||||
async def _restore_balance(
|
||||
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str
|
||||
session: AsyncSession,
|
||||
hashed_key: str,
|
||||
balance: int,
|
||||
reserved_balance: int,
|
||||
mint_url: str,
|
||||
) -> None:
|
||||
"""Restore balance after a failed refund mint attempt."""
|
||||
restore_stmt = (
|
||||
@@ -240,7 +363,11 @@ async def _restore_balance(
|
||||
await session.commit()
|
||||
logger.info(
|
||||
"refund_wallet_endpoint: balance restored after mint failure",
|
||||
extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url},
|
||||
extra={
|
||||
"hashed_key": hashed_key,
|
||||
"restored_balance": balance,
|
||||
"mint_url": mint_url,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -274,7 +401,20 @@ async def refund_wallet_endpoint(
|
||||
)
|
||||
out_tx = out_tx_result.first()
|
||||
if out_tx is None:
|
||||
raise HTTPException(status_code=404, detail="Refund not found")
|
||||
# The "in" row exists with a request_id, but the "out" (refund)
|
||||
# row hasn't been written yet — the upstream request is still in
|
||||
# flight and the refund will be minted once it completes. Tell the
|
||||
# client to retry instead of 404ing permanently (race condition
|
||||
# where /v1/wallet/refund is polled before the refund exists).
|
||||
logger.debug(
|
||||
"refund_wallet_endpoint: refund pending (in row exists, out row not yet created)",
|
||||
extra={"request_id": in_tx.request_id},
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=425,
|
||||
detail="Refund is pending; retry shortly.",
|
||||
headers={"Retry-After": "2"},
|
||||
)
|
||||
if out_tx.swept:
|
||||
raise HTTPException(status_code=410, detail="Refund has been swept")
|
||||
|
||||
@@ -305,6 +445,8 @@ async def refund_wallet_endpoint(
|
||||
if key.total_balance <= 0:
|
||||
if cached := await _refund_cache_get(bearer_value):
|
||||
return cached
|
||||
if persisted := await _get_persisted_api_key_refund(key, session):
|
||||
return persisted
|
||||
|
||||
if key.parent_key_hash:
|
||||
raise HTTPException(
|
||||
@@ -313,30 +455,19 @@ async def refund_wallet_endpoint(
|
||||
)
|
||||
|
||||
if key.reserved_balance > 0:
|
||||
# Release the reservation if it is stale
|
||||
cutoff = int(time.time()) - settings.stale_reservation_timeout_seconds
|
||||
stale_release_stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
||||
.where(col(ApiKey.reserved_balance) > 0)
|
||||
.where(
|
||||
or_(
|
||||
col(ApiKey.reserved_at).is_(None),
|
||||
col(ApiKey.reserved_at) < cutoff,
|
||||
)
|
||||
)
|
||||
.values(reserved_balance=0, reserved_at=None)
|
||||
# Release only durable reservations old enough to be stale. A newer
|
||||
# request on the same aggregate balance must remain reserved.
|
||||
await release_stale_reservations(
|
||||
session,
|
||||
settings.stale_reservation_timeout_seconds,
|
||||
key_hash=key.hashed_key,
|
||||
)
|
||||
stale_result = await session.exec(stale_release_stmt) # type: ignore[call-overload]
|
||||
await session.commit()
|
||||
|
||||
if stale_result.rowcount == 0:
|
||||
await session.refresh(key)
|
||||
if key.reserved_balance > 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot refund key. There are ongoing requests for this api key.",
|
||||
)
|
||||
|
||||
await session.refresh(key)
|
||||
logger.warning(
|
||||
"refund_wallet_endpoint: released stale reservation before refund",
|
||||
extra={
|
||||
@@ -381,15 +512,14 @@ async def refund_wallet_endpoint(
|
||||
detail="Balance changed concurrently. Please retry the refund.",
|
||||
)
|
||||
|
||||
# --- MINT: balance is locked at zero, safe to create the refund token ---
|
||||
# Proofs from untrusted mints are swapped to primary_mint on receive.
|
||||
# Use primary_mint unless key.refund_mint_url is an explicitly trusted mint.
|
||||
# The balance is locked at zero, so it is safe to create the refund token.
|
||||
effective_refund_mint = (
|
||||
key.refund_mint_url
|
||||
if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints
|
||||
else settings.primary_mint
|
||||
)
|
||||
try:
|
||||
refund_currency = key.refund_currency or "sat"
|
||||
if key.refund_address:
|
||||
await send_to_lnurl(
|
||||
remaining_balance,
|
||||
@@ -399,10 +529,10 @@ async def refund_wallet_endpoint(
|
||||
)
|
||||
result = {"recipient": key.refund_address}
|
||||
else:
|
||||
refund_currency = key.refund_currency or "sat"
|
||||
token = await send_token(
|
||||
remaining_balance, refund_currency, effective_refund_mint
|
||||
)
|
||||
effective_refund_mint = token_mint_url(token, effective_refund_mint)
|
||||
result = {"token": token}
|
||||
|
||||
if key.refund_currency == "sat":
|
||||
@@ -421,13 +551,47 @@ async def refund_wallet_endpoint(
|
||||
},
|
||||
)
|
||||
|
||||
except MeltOutcomeAmbiguousError as e:
|
||||
# The melt was dispatched and may still settle. Restoring the balance
|
||||
# here would let the same debit be paid out twice; keep the debit and
|
||||
# leave the outcome to reconciliation.
|
||||
logger.error(
|
||||
"refund_wallet_endpoint: melt outcome ambiguous; balance withheld "
|
||||
"pending reconciliation",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"hashed_key": key.hashed_key,
|
||||
"remaining_balance": remaining_balance,
|
||||
"refund_currency": key.refund_currency,
|
||||
"refund_mint_url": key.refund_mint_url,
|
||||
},
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail=(
|
||||
"Refund was dispatched but its outcome is unconfirmed; the "
|
||||
"balance is withheld until reconciliation completes"
|
||||
),
|
||||
)
|
||||
except HTTPException:
|
||||
# Minting failed — restore the debited balance
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
|
||||
await _restore_balance(
|
||||
session,
|
||||
key.hashed_key,
|
||||
pre_debit_balance,
|
||||
pre_debit_reserved,
|
||||
key.refund_mint_url or "",
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
# Minting failed — restore the debited balance
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
|
||||
await _restore_balance(
|
||||
session,
|
||||
key.hashed_key,
|
||||
pre_debit_balance,
|
||||
pre_debit_reserved,
|
||||
key.refund_mint_url or "",
|
||||
)
|
||||
error_msg = str(e)
|
||||
logger.error(
|
||||
"refund_wallet_endpoint: mint/send failed",
|
||||
@@ -441,14 +605,10 @@ async def refund_wallet_endpoint(
|
||||
"has_refund_address": bool(key.refund_address),
|
||||
},
|
||||
)
|
||||
if (
|
||||
"mint" in error_msg.lower()
|
||||
or "connection" in error_msg.lower()
|
||||
or "ConnectError" in str(type(e))
|
||||
):
|
||||
raise HTTPException(status_code=503, detail=f"Mint service unavailable: {error_msg}")
|
||||
if is_mint_connection_error(e):
|
||||
raise HTTPException(status_code=503, detail="Mint service unavailable")
|
||||
else:
|
||||
raise HTTPException(status_code=500, detail=f"Refund failed: {error_msg}")
|
||||
raise HTTPException(status_code=500, detail="Refund failed")
|
||||
|
||||
await _refund_cache_set(bearer_value, result)
|
||||
|
||||
@@ -458,7 +618,7 @@ async def refund_wallet_endpoint(
|
||||
token=result["token"],
|
||||
amount=remaining_balance,
|
||||
unit=key.refund_currency or "sat",
|
||||
mint_url=key.refund_mint_url,
|
||||
mint_url=effective_refund_mint,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="apikey",
|
||||
@@ -652,7 +812,6 @@ async def reset_child_key_spent(
|
||||
return {"success": True, "message": "Child key balance reset successfully."}
|
||||
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/{path:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE"],
|
||||
|
||||
+307
-71
@@ -13,13 +13,8 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from ..payment.models import _row_to_model, list_models
|
||||
from ..proxy import refresh_model_maps, reinitialize_upstreams
|
||||
from ..wallet import (
|
||||
fetch_all_balances,
|
||||
get_proofs_per_mint_and_unit,
|
||||
get_wallet,
|
||||
send_token,
|
||||
slow_filter_spend_proofs,
|
||||
)
|
||||
from ..wallet import fetch_all_balances, send_token, token_mint_url
|
||||
from . import vault
|
||||
from .db import (
|
||||
ApiKey,
|
||||
CashuTransaction,
|
||||
@@ -28,11 +23,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__)
|
||||
|
||||
@@ -44,6 +45,13 @@ ADMIN_SESSION_DURATION = 3600
|
||||
MAX_USAGE_ANALYTICS_HOURS = 365 * 24
|
||||
|
||||
|
||||
async def _refresh_provider_model_paths(upstream_provider_id: int) -> None:
|
||||
"""Queue discovery sync without blocking the committed admin mutation."""
|
||||
from ..upstream.model_paths import schedule_model_paths_refresh_for_provider
|
||||
|
||||
await schedule_model_paths_refresh_for_provider(upstream_provider_id)
|
||||
|
||||
|
||||
async def require_admin_api(request: Request) -> None:
|
||||
auth_header = request.headers.get("Authorization")
|
||||
if not auth_header or not auth_header.startswith("Bearer "):
|
||||
@@ -203,8 +211,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 +227,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 +245,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 +252,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 +319,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)
|
||||
@@ -408,33 +436,45 @@ class WithdrawRequest(BaseModel):
|
||||
async def withdraw(
|
||||
request: Request, withdraw_request: WithdrawRequest
|
||||
) -> dict[str, str]:
|
||||
# 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
|
||||
)
|
||||
proofs = get_proofs_per_mint_and_unit(
|
||||
wallet,
|
||||
withdraw_request.mint_url or global_settings.primary_mint,
|
||||
withdraw_request.unit,
|
||||
not_reserved=True,
|
||||
)
|
||||
proofs = await slow_filter_spend_proofs(proofs, wallet)
|
||||
current_balance = sum(proof.amount for proof in proofs)
|
||||
|
||||
effective_mint = withdraw_request.mint_url or global_settings.primary_mint
|
||||
if withdraw_request.amount <= 0:
|
||||
raise HTTPException(
|
||||
status_code=400, detail="Withdrawal amount must be positive"
|
||||
)
|
||||
|
||||
if withdraw_request.amount > current_balance:
|
||||
raise HTTPException(status_code=400, detail="Insufficient wallet balance")
|
||||
|
||||
token = await send_token(
|
||||
withdraw_request.amount, withdraw_request.unit, withdraw_request.mint_url
|
||||
)
|
||||
return {"token": token}
|
||||
try:
|
||||
token = await send_token(
|
||||
withdraw_request.amount, withdraw_request.unit, effective_mint
|
||||
)
|
||||
except ValueError as error:
|
||||
if not str(error).startswith("No trusted mint has "):
|
||||
raise
|
||||
raise HTTPException(
|
||||
status_code=400, detail="Insufficient wallet balance"
|
||||
) from error
|
||||
actual_mint = token_mint_url(token, effective_mint)
|
||||
try:
|
||||
await store_cashu_transaction(
|
||||
token=token,
|
||||
amount=withdraw_request.amount,
|
||||
unit=withdraw_request.unit,
|
||||
mint_url=actual_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": actual_mint,
|
||||
},
|
||||
)
|
||||
return {"token": token, "mint_url": actual_mint}
|
||||
|
||||
|
||||
class ModelCreate(BaseModel):
|
||||
@@ -461,7 +501,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}"
|
||||
)
|
||||
@@ -535,6 +574,7 @@ async def upsert_provider_model(
|
||||
await session.refresh(row)
|
||||
|
||||
await refresh_model_maps()
|
||||
await _refresh_provider_model_paths(provider_pk)
|
||||
return _row_to_model(
|
||||
row, apply_provider_fee=True, provider_fee=provider.provider_fee
|
||||
).dict() # type: ignore
|
||||
@@ -589,6 +629,7 @@ async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, ob
|
||||
await session.delete(row)
|
||||
await session.commit()
|
||||
await refresh_model_maps()
|
||||
await _refresh_provider_model_paths(provider_pk)
|
||||
return {"ok": True, "deleted_id": model_id}
|
||||
|
||||
|
||||
@@ -608,6 +649,7 @@ async def delete_all_provider_models(provider_id: str) -> dict[str, object]:
|
||||
await session.delete(row) # type: ignore
|
||||
await session.commit()
|
||||
await refresh_model_maps()
|
||||
await _refresh_provider_model_paths(provider_pk)
|
||||
return {"ok": True, "deleted": len(rows)}
|
||||
|
||||
|
||||
@@ -699,6 +741,7 @@ async def batch_override_provider_models(
|
||||
await session.commit()
|
||||
|
||||
await refresh_model_maps()
|
||||
await _refresh_provider_model_paths(provider_pk)
|
||||
return {
|
||||
"ok": True,
|
||||
"count": overridden_count,
|
||||
@@ -819,6 +862,33 @@ class UpstreamProviderUpdateBySlug(BaseModel):
|
||||
provider_settings: dict | None = None
|
||||
|
||||
|
||||
async def _active_ppq_claim_in_session(session: AsyncSession, provider_id: int) -> bool:
|
||||
"""Check for an active claim inside the caller's transaction.
|
||||
|
||||
Must share the transaction of whatever destructive write it is guarding —
|
||||
a check in its own session leaves a window for a worker to create the
|
||||
claim between the check and the commit.
|
||||
"""
|
||||
from ..upstream.auto_topup import _ppq_state_id_for_provider
|
||||
|
||||
claim = await session.get(CashuTransaction, _ppq_state_id_for_provider(provider_id))
|
||||
return claim is not None and not claim.collected and not claim.swept
|
||||
|
||||
|
||||
def _require_valid_ppq_auto_topup(
|
||||
provider_type: str, settings: dict | None
|
||||
) -> None:
|
||||
"""Reject PPQ auto top-up settings the worker would later refuse."""
|
||||
if provider_type != "ppqai":
|
||||
return
|
||||
|
||||
from ..upstream.auto_topup import validate_ppq_auto_topup_settings
|
||||
|
||||
problem = validate_ppq_auto_topup_settings(settings)
|
||||
if problem is not None:
|
||||
raise HTTPException(status_code=400, detail=problem)
|
||||
|
||||
|
||||
async def _apply_provider_update(
|
||||
session: AsyncSession,
|
||||
provider: UpstreamProviderRow,
|
||||
@@ -830,6 +900,29 @@ async def _apply_provider_update(
|
||||
await _ensure_unique_slug(session, validated, exclude_id=provider.id)
|
||||
provider.slug = validated
|
||||
|
||||
provider_type_changed = (
|
||||
payload.provider_type is not None
|
||||
and payload.provider_type != provider.provider_type
|
||||
)
|
||||
ppq_type_changed = provider_type_changed and (
|
||||
provider.provider_type == "ppqai" or payload.provider_type == "ppqai"
|
||||
)
|
||||
if (
|
||||
provider_type_changed
|
||||
and provider.provider_type == "ppqai"
|
||||
and provider.id is not None
|
||||
and await _active_ppq_claim_in_session(session, provider.id)
|
||||
):
|
||||
# Changing the type would orphan the claim: the PPQ endpoints refuse
|
||||
# non-ppqai providers, so nobody could ever inspect or release it.
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=(
|
||||
"This provider has an active PPQ auto top-up claim. Release "
|
||||
"it before changing the provider type"
|
||||
),
|
||||
)
|
||||
|
||||
if payload.provider_type is not None:
|
||||
provider.provider_type = payload.provider_type
|
||||
if payload.base_url is not None:
|
||||
@@ -842,6 +935,41 @@ async def _apply_provider_update(
|
||||
provider.enabled = payload.enabled
|
||||
if payload.provider_fee is not None:
|
||||
provider.provider_fee = payload.provider_fee
|
||||
|
||||
# Auto-top-up fields have provider-specific units and meaning. Reusing
|
||||
# enabled Routstr settings for PPQ (or vice versa) can silently reinterpret
|
||||
# sats as USD, so a type change must provide settings for the new type.
|
||||
if (
|
||||
ppq_type_changed
|
||||
and payload.provider_settings is None
|
||||
and provider.provider_settings
|
||||
):
|
||||
try:
|
||||
stored_settings = json.loads(provider.provider_settings)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
stored_settings = None
|
||||
if isinstance(stored_settings, dict) and stored_settings.get("auto_topup"):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
"Changing provider type requires explicit auto-top-up "
|
||||
"settings because the units are provider-specific"
|
||||
),
|
||||
)
|
||||
|
||||
# Validate against the effective type and effective settings.
|
||||
effective_settings = payload.provider_settings
|
||||
if effective_settings is None and payload.provider_type is not None:
|
||||
try:
|
||||
effective_settings = (
|
||||
json.loads(provider.provider_settings)
|
||||
if provider.provider_settings
|
||||
else None
|
||||
)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
effective_settings = None
|
||||
if effective_settings is not None:
|
||||
_require_valid_ppq_auto_topup(provider.provider_type, effective_settings)
|
||||
if payload.provider_settings is not None:
|
||||
provider.provider_settings = json.dumps(payload.provider_settings)
|
||||
|
||||
@@ -881,6 +1009,10 @@ async def create_upstream_provider(
|
||||
else:
|
||||
slug = await allocate_unique_provider_slug(session, payload.provider_type)
|
||||
|
||||
_require_valid_ppq_auto_topup(
|
||||
payload.provider_type, payload.provider_settings
|
||||
)
|
||||
|
||||
provider = UpstreamProviderRow(
|
||||
slug=slug,
|
||||
provider_type=payload.provider_type,
|
||||
@@ -899,6 +1031,7 @@ async def create_upstream_provider(
|
||||
|
||||
await reinitialize_upstreams()
|
||||
await refresh_model_maps()
|
||||
await _refresh_provider_model_paths(_provider_pk(provider))
|
||||
return _serialize_provider(provider)
|
||||
|
||||
|
||||
@@ -924,6 +1057,7 @@ async def update_upstream_provider(
|
||||
|
||||
await reinitialize_upstreams()
|
||||
await refresh_model_maps()
|
||||
await _refresh_provider_model_paths(_provider_pk(provider))
|
||||
return _serialize_provider(provider)
|
||||
|
||||
|
||||
@@ -959,6 +1093,7 @@ async def update_upstream_provider_by_slug(
|
||||
|
||||
await reinitialize_upstreams()
|
||||
await refresh_model_maps()
|
||||
await _refresh_provider_model_paths(_provider_pk(provider))
|
||||
return _serialize_provider(provider)
|
||||
|
||||
|
||||
@@ -969,6 +1104,25 @@ async def delete_upstream_provider(provider_id: str) -> dict[str, object]:
|
||||
async with create_session() as session:
|
||||
provider = await _get_upstream_provider_by_ref(session, provider_id)
|
||||
deleted_id = _provider_pk(provider)
|
||||
|
||||
# Checked inside the delete transaction: the worker's claim creation
|
||||
# re-reads the provider inside its own transaction, so these two
|
||||
# writes serialise — either the claim lands first and this 409s, or
|
||||
# the delete lands first and the worker refuses to claim.
|
||||
if provider.provider_type == "ppqai" and await _active_ppq_claim_in_session(
|
||||
session, deleted_id
|
||||
):
|
||||
# Deleting now would orphan the claim and any funds it tracks:
|
||||
# the PPQ endpoints 404 without the provider row, so the claim
|
||||
# could never again be inspected or released.
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=(
|
||||
"This provider has an active PPQ auto top-up claim. "
|
||||
"Resolve and release it before deleting the provider"
|
||||
),
|
||||
)
|
||||
|
||||
await session.delete(provider)
|
||||
await session.commit()
|
||||
await reinitialize_upstreams()
|
||||
@@ -1577,6 +1731,78 @@ async def get_log_dates_api(request: Request) -> dict[str, object]:
|
||||
return {"dates": dates}
|
||||
|
||||
|
||||
_PPQ_RELEASE_ERRORS = {
|
||||
"no_active_claim": "No active PPQ claim to release",
|
||||
"stale_state": ("The claim changed since it was reviewed; reload and check again"),
|
||||
"payment_in_flight": (
|
||||
"A Lightning payment is still in flight for this claim. Wait for it to "
|
||||
"finish or expire before releasing"
|
||||
),
|
||||
"claim_changed": (
|
||||
"The claim changed while the release was being applied; reload and check again"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
class ReleasePPQAutoTopupRequest(BaseModel):
|
||||
confirmed_safe_to_retry: bool
|
||||
# Echoes the state_token the admin reviewed — the claim's full versioned
|
||||
# state, not just its operation id. Any change since the review (a new
|
||||
# attempt, a phase change, a renewed lease) fails the match, so the
|
||||
# release cannot land on a state the admin never saw.
|
||||
state_token: str | None = None
|
||||
|
||||
|
||||
async def _require_ppq_provider(provider_id: int) -> UpstreamProviderRow:
|
||||
async with create_session() as session:
|
||||
provider = await session.get(UpstreamProviderRow, provider_id)
|
||||
if provider is None:
|
||||
raise HTTPException(status_code=404, detail="Provider not found")
|
||||
if provider.provider_type != "ppqai":
|
||||
raise HTTPException(status_code=400, detail="Provider is not PPQ")
|
||||
return provider
|
||||
|
||||
|
||||
@admin_router.get(
|
||||
"/api/upstream-providers/{provider_id}/ppq-auto-topup",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def get_ppq_auto_topup_api(provider_id: int) -> dict[str, object]:
|
||||
await _require_ppq_provider(provider_id)
|
||||
from ..upstream.auto_topup import get_ppq_auto_topup_state
|
||||
|
||||
return {"ok": True, **await get_ppq_auto_topup_state(provider_id)}
|
||||
|
||||
|
||||
@admin_router.post(
|
||||
"/api/upstream-providers/{provider_id}/ppq-auto-topup/release",
|
||||
dependencies=[Depends(require_admin_api)],
|
||||
)
|
||||
async def release_ppq_auto_topup_api(
|
||||
provider_id: int, payload: ReleasePPQAutoTopupRequest
|
||||
) -> dict[str, object]:
|
||||
await _require_ppq_provider(provider_id)
|
||||
if not payload.confirmed_safe_to_retry:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Confirm the Lightning payment outcome is safe before releasing",
|
||||
)
|
||||
|
||||
from ..upstream.auto_topup import release_ppq_auto_topup_state
|
||||
|
||||
outcome = await release_ppq_auto_topup_state(
|
||||
provider_id, state_token=payload.state_token
|
||||
)
|
||||
if not outcome.released:
|
||||
raise HTTPException(status_code=409, detail=_PPQ_RELEASE_ERRORS[outcome.reason])
|
||||
|
||||
logger.warning(
|
||||
"Admin released PPQ auto top-up claim after manual reconciliation",
|
||||
extra={"provider_id": provider_id, "state_token": payload.state_token},
|
||||
)
|
||||
return {"ok": True, "released": True}
|
||||
|
||||
|
||||
@admin_router.get("/api/transactions", dependencies=[Depends(require_admin_api)])
|
||||
async def get_transactions_api(
|
||||
type: str | None = None,
|
||||
@@ -1589,7 +1815,11 @@ async def get_transactions_api(
|
||||
async with create_session() as session:
|
||||
from sqlmodel import col, func
|
||||
|
||||
base = select(CashuTransaction)
|
||||
# Hide only the deterministic PPQ claim-lock rows. Append-only PPQ
|
||||
# payment rows remain visible as the audit trail for irreversible melts.
|
||||
base = select(CashuTransaction).where(
|
||||
~col(CashuTransaction.id).like("ppq-auto-topup-%")
|
||||
)
|
||||
if type:
|
||||
base = base.where(CashuTransaction.type == type)
|
||||
if source:
|
||||
@@ -1625,12 +1855,18 @@ async def get_transactions_api(
|
||||
)
|
||||
total = count_result.one()
|
||||
|
||||
stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit)
|
||||
stmt = (
|
||||
base.order_by(col(CashuTransaction.created_at).desc())
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
)
|
||||
results = await session.exec(stmt)
|
||||
transactions = results.all()
|
||||
|
||||
return {
|
||||
"transactions": [tx.dict() for tx in transactions],
|
||||
"transactions": [
|
||||
tx.dict(exclude={"sweep_started_at"}) for tx in transactions
|
||||
],
|
||||
"total": total,
|
||||
}
|
||||
|
||||
|
||||
+508
-39
@@ -1,29 +1,91 @@
|
||||
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 import Index, UniqueConstraint, case, delete, event, or_
|
||||
from sqlalchemy.engine import make_url
|
||||
from sqlalchemy.exc import IntegrityError, OperationalError
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlalchemy.ext.asyncio.engine import create_async_engine
|
||||
from sqlalchemy.orm import aliased
|
||||
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from .logging import get_logger
|
||||
from .settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db")
|
||||
|
||||
|
||||
engine = create_async_engine(DATABASE_URL, echo=False) # echo=True for debugging SQL
|
||||
def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
|
||||
"""Build and instrument an async engine from environment-only settings."""
|
||||
url = make_url(database_url)
|
||||
backend = url.get_backend_name()
|
||||
is_sqlite = backend == "sqlite"
|
||||
is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"}
|
||||
pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite
|
||||
options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping}
|
||||
if not is_memory_sqlite:
|
||||
options.update(
|
||||
pool_size=settings.database_pool_size,
|
||||
max_overflow=settings.database_max_overflow,
|
||||
pool_timeout=settings.database_pool_timeout,
|
||||
pool_recycle=settings.database_pool_recycle,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Database pool configured",
|
||||
extra={
|
||||
"database_url_backend": backend,
|
||||
"in_memory_sqlite": is_memory_sqlite,
|
||||
**options,
|
||||
},
|
||||
)
|
||||
created_engine = create_async_engine(database_url, echo=False, **options)
|
||||
hold_warn_seconds = settings.database_pool_hold_warn_seconds
|
||||
|
||||
def record_pool_checkout(
|
||||
dbapi_connection: object, connection_record: object, proxy: object
|
||||
) -> None:
|
||||
connection_record.info["routstr_checked_out_at"] = time.monotonic() # type: ignore[attr-defined]
|
||||
|
||||
def record_pool_checkin(
|
||||
dbapi_connection: object, connection_record: object
|
||||
) -> None:
|
||||
checked_out_at = connection_record.info.pop( # type: ignore[attr-defined]
|
||||
"routstr_checked_out_at", None
|
||||
)
|
||||
if checked_out_at is None:
|
||||
return
|
||||
held_seconds = time.monotonic() - checked_out_at
|
||||
if held_seconds >= hold_warn_seconds:
|
||||
logger.warning(
|
||||
"Database connection held longer than threshold",
|
||||
extra={
|
||||
"held_seconds": round(held_seconds, 3),
|
||||
"threshold_seconds": hold_warn_seconds,
|
||||
"pool_status": created_engine.pool.status(),
|
||||
},
|
||||
)
|
||||
|
||||
event.listen(created_engine.sync_engine, "checkout", record_pool_checkout)
|
||||
event.listen(created_engine.sync_engine, "checkin", record_pool_checkin)
|
||||
return created_engine
|
||||
|
||||
|
||||
engine = create_db_engine()
|
||||
|
||||
|
||||
class ApiKey(SQLModel, table=True): # type: ignore
|
||||
@@ -96,32 +158,133 @@ class ApiKey(SQLModel, table=True): # type: ignore
|
||||
|
||||
|
||||
async def reset_all_reserved_balances(session: AsyncSession) -> None:
|
||||
stmt = update(ApiKey).values(reserved_balance=0, reserved_at=None)
|
||||
await session.exec(stmt) # type: ignore[call-overload]
|
||||
"""Release every active durable reservation during explicit startup reset."""
|
||||
await session.exec( # type: ignore[call-overload]
|
||||
update(ReservationRelease)
|
||||
.where(col(ReservationRelease.status) == "active")
|
||||
.values(status="released")
|
||||
)
|
||||
await session.exec( # type: ignore[call-overload]
|
||||
update(ApiKey).values(reserved_balance=0, reserved_at=None)
|
||||
)
|
||||
await session.commit()
|
||||
logger.info("Reset reserved balances on startup")
|
||||
|
||||
|
||||
async def release_stale_reservations(
|
||||
session: AsyncSession, max_age_seconds: int
|
||||
session: AsyncSession,
|
||||
max_age_seconds: int,
|
||||
*,
|
||||
key_hash: str | None = None,
|
||||
) -> int:
|
||||
"""Release reservations whose last reserve is older than max_age_seconds.
|
||||
"""
|
||||
"""Release stale durable reservations without touching newer reservations."""
|
||||
cutoff = int(time.time()) - max_age_seconds
|
||||
stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.reserved_balance) > 0)
|
||||
.where(col(ApiKey.reserved_at).is_not(None))
|
||||
.where(col(ApiKey.reserved_at) < cutoff)
|
||||
.values(reserved_balance=0, reserved_at=None)
|
||||
query = (
|
||||
select(ReservationRelease)
|
||||
.where(col(ReservationRelease.status) == "active")
|
||||
.where(col(ReservationRelease.created_at) < cutoff)
|
||||
)
|
||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
||||
if key_hash is not None:
|
||||
query = query.where(
|
||||
or_(
|
||||
col(ReservationRelease.key_hash) == key_hash,
|
||||
col(ReservationRelease.billing_key_hash) == key_hash,
|
||||
)
|
||||
)
|
||||
reservations = (await session.exec(query)).all()
|
||||
released = 0
|
||||
|
||||
for reservation in reservations:
|
||||
transition = await session.exec( # type: ignore[call-overload]
|
||||
update(ReservationRelease)
|
||||
.where(col(ReservationRelease.id) == reservation.id)
|
||||
.where(col(ReservationRelease.status) == "active")
|
||||
.values(status="released")
|
||||
)
|
||||
if transition.rowcount != 1:
|
||||
continue
|
||||
|
||||
values = {
|
||||
"reserved_balance": col(ApiKey.reserved_balance)
|
||||
- reservation.reserved_msats,
|
||||
"reserved_at": case(
|
||||
(
|
||||
col(ApiKey.reserved_balance) - reservation.reserved_msats > 0,
|
||||
col(ApiKey.reserved_at),
|
||||
),
|
||||
else_=None,
|
||||
),
|
||||
}
|
||||
parent_result = await session.exec( # type: ignore[call-overload]
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == reservation.billing_key_hash)
|
||||
.where(col(ApiKey.reserved_balance) >= reservation.reserved_msats)
|
||||
.values(**values)
|
||||
)
|
||||
if parent_result.rowcount != 1:
|
||||
await session.rollback()
|
||||
return 0
|
||||
|
||||
if reservation.billing_key_hash != reservation.key_hash:
|
||||
child_result = await session.exec( # type: ignore[call-overload]
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == reservation.key_hash)
|
||||
.where(col(ApiKey.reserved_balance) >= reservation.reserved_msats)
|
||||
.values(**values)
|
||||
)
|
||||
if child_result.rowcount != 1:
|
||||
await session.rollback()
|
||||
return 0
|
||||
released += 1
|
||||
|
||||
# Rolling upgrades can leave aggregate reservations created before durable
|
||||
# reservation rows existed. Release only stale aggregates that have no active
|
||||
# durable owner; targeted refund cleanup also heals legacy NULL timestamps.
|
||||
legacy_query = select(ApiKey).where(col(ApiKey.reserved_balance) > 0)
|
||||
if key_hash is None:
|
||||
legacy_query = legacy_query.where(col(ApiKey.reserved_at).is_not(None)).where(
|
||||
col(ApiKey.reserved_at) < cutoff
|
||||
)
|
||||
else:
|
||||
legacy_query = legacy_query.where(
|
||||
or_(
|
||||
col(ApiKey.hashed_key) == key_hash,
|
||||
col(ApiKey.parent_key_hash) == key_hash,
|
||||
)
|
||||
).where(
|
||||
or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff)
|
||||
)
|
||||
|
||||
for legacy_key in (await session.exec(legacy_query)).all():
|
||||
active_owner = (
|
||||
await session.exec(
|
||||
select(ReservationRelease.id)
|
||||
.where(col(ReservationRelease.status) == "active")
|
||||
.where(
|
||||
or_(
|
||||
col(ReservationRelease.key_hash) == legacy_key.hashed_key,
|
||||
col(ReservationRelease.billing_key_hash)
|
||||
== legacy_key.hashed_key,
|
||||
)
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
).first()
|
||||
if active_owner is not None:
|
||||
continue
|
||||
legacy_key.reserved_balance = 0
|
||||
legacy_key.reserved_at = None
|
||||
session.add(legacy_key)
|
||||
released += 1
|
||||
|
||||
await session.commit()
|
||||
released = int(result.rowcount or 0)
|
||||
if released:
|
||||
logger.warning(
|
||||
"Released stale balance reservations",
|
||||
extra={"released_keys": released, "max_age_seconds": max_age_seconds},
|
||||
"Released stale reservations",
|
||||
extra={
|
||||
"released_reservations": released,
|
||||
"max_age_seconds": max_age_seconds,
|
||||
},
|
||||
)
|
||||
return released
|
||||
|
||||
@@ -130,7 +293,7 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
|
||||
"""Delete dead parentless API keys; return the count removed.
|
||||
|
||||
Dead = 0 balance/reservation/spend/requests, older than the grace period,
|
||||
no parent, no children, no pending invoice. Cashu rows are unlinked (not
|
||||
no parent, no children, no retryable invoice. Cashu rows are unlinked (not
|
||||
deleted) first to keep the audit trail.
|
||||
"""
|
||||
cutoff = int(time.time()) - min_age_seconds
|
||||
@@ -144,7 +307,9 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
|
||||
pending_invoice = (
|
||||
select(LightningInvoice.id)
|
||||
.where(col(LightningInvoice.api_key_hash) == col(ApiKey.hashed_key))
|
||||
.where(col(LightningInvoice.status) == "pending")
|
||||
.where(
|
||||
col(LightningInvoice.status).in_(("pending", "settlement_pending"))
|
||||
)
|
||||
).exists()
|
||||
|
||||
eligible_hashes = (
|
||||
@@ -154,9 +319,7 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
|
||||
.where(col(ApiKey.total_spent) == 0)
|
||||
.where(col(ApiKey.total_requests) == 0)
|
||||
.where(col(ApiKey.parent_key_hash).is_(None))
|
||||
.where(
|
||||
(col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff)
|
||||
)
|
||||
.where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff))
|
||||
.where(~pending_invoice)
|
||||
.where(~has_children)
|
||||
)
|
||||
@@ -210,6 +373,60 @@ class ModelRow(SQLModel, table=True): # type: ignore
|
||||
upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models")
|
||||
|
||||
|
||||
class ModelPathRow(SQLModel, table=True): # type: ignore
|
||||
"""Upstream provider path a model is reachable through.
|
||||
|
||||
Discovery/visibility data only. ``model_id`` is intentionally NOT globally
|
||||
unique: it is the client-visible ``/v1/models`` id (``forwarded_model_id or
|
||||
id``) grouped across every provider that exposes the model. A single model
|
||||
can therefore have several rows — one per direct provider path plus one per
|
||||
OpenRouter sub-provider endpoint.
|
||||
"""
|
||||
|
||||
__tablename__ = "model_paths"
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"model_id",
|
||||
"path",
|
||||
"upstream_provider_id",
|
||||
name="uq_model_paths_model_path_provider",
|
||||
),
|
||||
)
|
||||
id: int | None = Field(default=None, primary_key=True)
|
||||
# No standalone index on model_id: the unique constraint's autoindex already
|
||||
# leads on model_id, so a second index only adds write amplification.
|
||||
model_id: str = Field(
|
||||
description="Client-visible /v1/models id (forwarded_model_id or id)"
|
||||
)
|
||||
path: str = Field(
|
||||
description=(
|
||||
"Opaque selector containing upstream URL, provider ID, model ID, "
|
||||
"and optional endpoint tag"
|
||||
)
|
||||
)
|
||||
provider_slug: str = Field(
|
||||
description="Public slug of the configured upstream provider"
|
||||
)
|
||||
provider_type: str = Field(description="Configured upstream provider type")
|
||||
endpoint_tag: str | None = Field(
|
||||
default=None,
|
||||
description="Exact OpenRouter endpoint tag used for request-side selection",
|
||||
)
|
||||
endpoint_name: str | None = Field(
|
||||
default=None, description="Human-readable endpoint display name"
|
||||
)
|
||||
upstream_provider_id: int = Field(
|
||||
index=True,
|
||||
foreign_key="upstream_providers.id",
|
||||
ondelete="CASCADE",
|
||||
description="upstream_providers.id this path was discovered from",
|
||||
)
|
||||
updated_at: int = Field(
|
||||
default=0,
|
||||
description="Unix timestamp of the refresh cycle that wrote this row",
|
||||
)
|
||||
|
||||
|
||||
class LightningInvoice(SQLModel, table=True): # type: ignore
|
||||
__tablename__ = "lightning_invoices"
|
||||
|
||||
@@ -219,12 +436,19 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
|
||||
description: str = Field(description="Invoice description")
|
||||
payment_hash: str = Field(description="Payment hash for tracking", unique=True)
|
||||
status: str = Field(
|
||||
default="pending", description="pending, paid, expired, cancelled"
|
||||
default="pending",
|
||||
description=(
|
||||
"pending, settlement_pending, paid, expired, cancelled, "
|
||||
"reconciliation_required"
|
||||
),
|
||||
)
|
||||
api_key_hash: str | None = Field(
|
||||
default=None, description="Associated API key hash for topup operations"
|
||||
)
|
||||
purpose: str = Field(description="create or topup")
|
||||
mint_url: str | None = Field(
|
||||
default=None, description="Mint URL where the quote was created (fallback tracking)"
|
||||
)
|
||||
created_at: int = Field(
|
||||
default_factory=lambda: int(time.time()), description="Unix timestamp"
|
||||
)
|
||||
@@ -264,6 +488,10 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
|
||||
)
|
||||
collected: bool = Field(default=False)
|
||||
swept: bool = Field(default=False)
|
||||
sweep_started_at: int | None = Field(
|
||||
default=None,
|
||||
description="Unix timestamp for a recoverable refund-sweep claim",
|
||||
)
|
||||
source: str = Field(
|
||||
default="x-cashu",
|
||||
description="Payment source: x-cashu or apikey",
|
||||
@@ -287,10 +515,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 +535,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
|
||||
@@ -346,21 +659,71 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
|
||||
)
|
||||
|
||||
|
||||
class ReservationRelease(SQLModel, table=True): # type: ignore
|
||||
__tablename__ = "reservation_releases"
|
||||
__table_args__ = (
|
||||
Index("ix_reservation_releases_status_created_at", "status", "created_at"),
|
||||
)
|
||||
|
||||
id: str = Field(primary_key=True)
|
||||
key_hash: str = Field(index=True)
|
||||
billing_key_hash: str = Field(index=True)
|
||||
reserved_msats: int
|
||||
status: str = Field(default="active")
|
||||
created_at: int = Field(default_factory=lambda: int(time.time()))
|
||||
|
||||
|
||||
class RoutstrFee(SQLModel, table=True): # type: ignore
|
||||
__tablename__ = "routstr_fees"
|
||||
id: int = Field(default=1, primary_key=True)
|
||||
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
|
||||
"""Long-lived authorization token for CLI/agent use against admin endpoints."""
|
||||
|
||||
__tablename__ = "cli_tokens"
|
||||
id: str = Field(
|
||||
primary_key=True, default_factory=lambda: uuid.uuid4().hex
|
||||
)
|
||||
id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex)
|
||||
token: str = Field(unique=True, index=True, description="Bearer token value")
|
||||
name: str = Field(description="Human-readable label for this token")
|
||||
created_at: int = Field(default_factory=lambda: int(time.time()))
|
||||
@@ -392,28 +755,134 @@ 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(
|
||||
async def total_user_liability(db_session: AsyncSession) -> int:
|
||||
"""Return all outstanding API-key balances in millisatoshis."""
|
||||
result = await db_session.exec(select(func.sum(ApiKey.balance)))
|
||||
return int(result.one() or 0)
|
||||
|
||||
|
||||
async def balance_for_mint_and_unit(
|
||||
db_session: AsyncSession, mint_url: str, unit: str
|
||||
) -> int:
|
||||
query = select(func.sum(ApiKey.balance)).where(
|
||||
ApiKey.refund_mint_url == mint_url, ApiKey.refund_currency == unit
|
||||
"""Return the user liability for one mint and unit in millisatoshis."""
|
||||
result = await db_session.exec(
|
||||
select(func.sum(ApiKey.balance)).where(
|
||||
col(ApiKey.refund_mint_url) == mint_url,
|
||||
col(ApiKey.refund_currency) == unit,
|
||||
)
|
||||
)
|
||||
return int(result.one() or 0)
|
||||
|
||||
|
||||
async def balances_by_mint_and_unit(
|
||||
db_session: AsyncSession, mint_urls: list[str], units: list[str]
|
||||
) -> dict[tuple[str, str], int]:
|
||||
"""Return requested user liabilities grouped by mint and unit."""
|
||||
if not mint_urls or not units:
|
||||
return {}
|
||||
query = (
|
||||
select(
|
||||
col(ApiKey.refund_mint_url),
|
||||
col(ApiKey.refund_currency),
|
||||
func.sum(ApiKey.balance),
|
||||
)
|
||||
.where(
|
||||
col(ApiKey.refund_mint_url).in_(mint_urls),
|
||||
col(ApiKey.refund_currency).in_(units),
|
||||
)
|
||||
.group_by(col(ApiKey.refund_mint_url), col(ApiKey.refund_currency))
|
||||
)
|
||||
result = await db_session.exec(query)
|
||||
return result.one() or 0
|
||||
return {
|
||||
(mint_url, unit): int(balance or 0)
|
||||
for mint_url, unit, balance in result.all()
|
||||
if mint_url is not None and unit is not None
|
||||
}
|
||||
|
||||
|
||||
async def init_db() -> None:
|
||||
|
||||
+24
-15
@@ -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__
|
||||
|
||||
@@ -58,6 +58,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
providers_task = None
|
||||
models_refresh_task = None
|
||||
model_maps_refresh_task = None
|
||||
model_paths_refresh_task = None
|
||||
key_reset_task = None
|
||||
stale_reservation_task = None
|
||||
dead_key_prune_task = None
|
||||
@@ -85,17 +86,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
|
||||
@@ -127,6 +131,13 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
refresh_upstreams_models_periodically(get_upstreams)
|
||||
)
|
||||
model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically())
|
||||
# Always started: the loop re-reads the enable flag and interval every
|
||||
# iteration, so 0 -> N (or re-enabling) takes effect without a restart.
|
||||
from ..upstream.model_paths import refresh_model_paths_periodically
|
||||
|
||||
model_paths_refresh_task = asyncio.create_task(
|
||||
refresh_model_paths_periodically(get_upstreams)
|
||||
)
|
||||
payout_task = asyncio.create_task(periodic_payout())
|
||||
if global_settings.nsec:
|
||||
nip91_task = asyncio.create_task(announce_provider())
|
||||
@@ -134,9 +145,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
if global_settings.providers_refresh_interval_seconds > 0:
|
||||
providers_task = asyncio.create_task(providers_cache_refresher())
|
||||
key_reset_task = asyncio.create_task(periodic_key_reset())
|
||||
stale_reservation_task = asyncio.create_task(
|
||||
periodic_stale_reservation_sweep()
|
||||
)
|
||||
stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep())
|
||||
dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune())
|
||||
auto_topup_task = asyncio.create_task(periodic_auto_topup())
|
||||
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
|
||||
@@ -173,6 +182,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
models_refresh_task.cancel()
|
||||
if model_maps_refresh_task is not None:
|
||||
model_maps_refresh_task.cancel()
|
||||
if model_paths_refresh_task is not None:
|
||||
model_paths_refresh_task.cancel()
|
||||
if key_reset_task is not None:
|
||||
key_reset_task.cancel()
|
||||
if stale_reservation_task is not None:
|
||||
@@ -206,6 +217,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
tasks_to_wait.append(models_refresh_task)
|
||||
if model_maps_refresh_task is not None:
|
||||
tasks_to_wait.append(model_maps_refresh_task)
|
||||
if model_paths_refresh_task is not None:
|
||||
tasks_to_wait.append(model_paths_refresh_task)
|
||||
if key_reset_task is not None:
|
||||
tasks_to_wait.append(key_reset_task)
|
||||
if stale_reservation_task is not None:
|
||||
@@ -242,9 +255,7 @@ class _ImmutableStaticFiles(StaticFiles):
|
||||
async def get_response(self, path: str, scope: Scope) -> StarletteResponse:
|
||||
response = await super().get_response(path, scope)
|
||||
if response.status_code == 200:
|
||||
response.headers["Cache-Control"] = (
|
||||
"public, max-age=31536000, immutable"
|
||||
)
|
||||
response.headers["Cache-Control"] = "public, max-age=31536000, immutable"
|
||||
return response
|
||||
|
||||
|
||||
@@ -318,9 +329,7 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
|
||||
# Serve the App Router RSC payload for the home page.
|
||||
@app.get("/index.txt", include_in_schema=False)
|
||||
async def serve_root_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "index.txt", media_type="text/x-component"
|
||||
)
|
||||
return FileResponse(UI_DIST_PATH / "index.txt", media_type="text/x-component")
|
||||
|
||||
# Next.js is built with `trailingSlash: true`, so all UI page URLs end
|
||||
# with a slash (e.g. `/login/`). The proxy router catches `/{path:path}`
|
||||
|
||||
+317
-32
@@ -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")
|
||||
@@ -40,6 +41,9 @@ class Settings(BaseSettings):
|
||||
receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS")
|
||||
primary_mint: str = Field(default="", env="PRIMARY_MINT_URL")
|
||||
primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT")
|
||||
mint_operation_concurrency: int = Field(
|
||||
default=4, ge=1, env="MINT_OPERATION_CONCURRENCY"
|
||||
)
|
||||
|
||||
# Lightning payout configuration
|
||||
# Minimum available balance (in satoshis) before profit is paid out over
|
||||
@@ -49,6 +53,18 @@ class Settings(BaseSettings):
|
||||
payout_interval_seconds: int = Field(
|
||||
default=900, gt=0, env="PAYOUT_INTERVAL_SECONDS"
|
||||
)
|
||||
# Timeout (seconds) for individual mint API operations (melt, mint, swap,
|
||||
# checkstate). When a mint is slow or rate-limiting, operations are
|
||||
# cancelled after this delay instead of hanging indefinitely.
|
||||
mint_operation_timeout_seconds: int = Field(
|
||||
default=30, gt=0, env="MINT_OPERATION_TIMEOUT_SECONDS"
|
||||
)
|
||||
# Maximum concurrent API operations per mint. Actual mint quotas vary by
|
||||
# endpoint, so 429 responses drive adaptive cooldown instead of fixed RPM
|
||||
# pacing. 0 = unlimited concurrency.
|
||||
mint_max_concurrency: int = Field(default=4, ge=0, env="MINT_MAX_CONCURRENCY")
|
||||
# Max retries when a mint returns 429 or times out (exponential backoff).
|
||||
mint_retry_max_attempts: int = Field(default=3, ge=0, env="MINT_RETRY_MAX_ATTEMPTS")
|
||||
|
||||
# Pricing
|
||||
# Default behavior: derive pricing from MODELS
|
||||
@@ -94,10 +110,36 @@ class Settings(BaseSettings):
|
||||
models_refresh_interval_seconds: int = Field(
|
||||
default=360, env="MODELS_REFRESH_INTERVAL_SECONDS"
|
||||
)
|
||||
model_paths_refresh_interval_seconds: int = Field(
|
||||
default=600, env="MODEL_PATHS_REFRESH_INTERVAL_SECONDS"
|
||||
)
|
||||
enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH")
|
||||
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
|
||||
enable_model_paths_refresh: bool = Field(
|
||||
default=True, env="ENABLE_MODEL_PATHS_REFRESH"
|
||||
)
|
||||
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
|
||||
refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS")
|
||||
refund_sweep_ttl_seconds: int = Field(
|
||||
default=604800, env="REFUND_SWEEP_TTL_SECONDS"
|
||||
)
|
||||
refund_sweep_claim_timeout_seconds: int = Field(
|
||||
default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS"
|
||||
)
|
||||
|
||||
# Database connection-pool controls (advanced). Capacity defaults provide
|
||||
# headroom for Routstr's concurrent request and background-payment workload.
|
||||
# Pre-ping is enabled by the engine factory for networked backends; SQLite
|
||||
# can explicitly opt in. These fields are env-only below.
|
||||
database_pool_size: int = Field(default=10, ge=1, env="DATABASE_POOL_SIZE")
|
||||
database_max_overflow: int = Field(default=20, ge=0, env="DATABASE_MAX_OVERFLOW")
|
||||
database_pool_timeout: float = Field(
|
||||
default=15.0, gt=0, env="DATABASE_POOL_TIMEOUT"
|
||||
)
|
||||
database_pool_recycle: int = Field(default=1800, ge=0, env="DATABASE_POOL_RECYCLE")
|
||||
database_pool_pre_ping: bool = Field(default=False, env="DATABASE_POOL_PRE_PING")
|
||||
database_pool_hold_warn_seconds: float = Field(
|
||||
default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS"
|
||||
)
|
||||
|
||||
# Logging
|
||||
log_level: str = Field(default="INFO", env="LOG_LEVEL")
|
||||
@@ -116,9 +158,8 @@ class Settings(BaseSettings):
|
||||
|
||||
# Discovery
|
||||
relays: list[str] = Field(default_factory=list, env="RELAYS")
|
||||
enable_analytics_sharing: bool = Field(
|
||||
default=True, env="ENABLE_ANALYTICS_SHARING"
|
||||
)
|
||||
enable_analytics_sharing: bool = Field(default=True, env="ENABLE_ANALYTICS_SHARING")
|
||||
|
||||
|
||||
def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Discard unknown keys from persisted settings."""
|
||||
@@ -132,10 +173,92 @@ 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"})
|
||||
|
||||
# Infrastructure the node needs *before* it can open a DB session — so it can
|
||||
# never be configured from the DB (chicken-and-egg) and stays env-only. Unlike
|
||||
# secrets (owned by bootstrap), these are excluded so the DB settings blob can
|
||||
# neither store nor shadow them; env is always authoritative.
|
||||
ENV_ONLY_FIELDS = frozenset(
|
||||
{
|
||||
"database_pool_size",
|
||||
"database_max_overflow",
|
||||
"database_pool_timeout",
|
||||
"database_pool_recycle",
|
||||
"database_pool_pre_ping",
|
||||
"database_pool_hold_warn_seconds",
|
||||
}
|
||||
)
|
||||
|
||||
_NON_PERSISTED_FIELDS = SECRET_FIELDS | ENV_ONLY_FIELDS
|
||||
|
||||
|
||||
def _strip_secret_fields(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Return a copy of ``data`` without secret or env-only fields.
|
||||
|
||||
Both are kept out of the persisted settings blob: secrets for confidentiality,
|
||||
env-only fields (e.g. DB pool sizing) because they must never be sourced from
|
||||
the database.
|
||||
"""
|
||||
return {k: v for k, v in data.items() if k not in _NON_PERSISTED_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 +313,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 +365,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
|
||||
@@ -281,7 +389,13 @@ class SettingsService:
|
||||
valid_fields = set(env_resolved.dict().keys())
|
||||
merged_dict: dict[str, Any] = dict(env_resolved.dict())
|
||||
merged_dict.update(
|
||||
{k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields}
|
||||
{
|
||||
k: v
|
||||
for k, v in db_json.items()
|
||||
if v not in (None, "", [], {})
|
||||
and k in valid_fields
|
||||
and k not in ENV_ONLY_FIELDS
|
||||
}
|
||||
)
|
||||
merged_dict = Settings(**merged_dict).dict()
|
||||
|
||||
@@ -291,20 +405,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,13 +457,18 @@ 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),
|
||||
)
|
||||
)
|
||||
await db_session.commit()
|
||||
# Update in-place
|
||||
# Update in-place. Env-only fields (e.g. DB pool sizing) are never
|
||||
# applied here: the engine pool is already built at boot from env,
|
||||
# so letting an update mutate the live value would only make it
|
||||
# diverge from the running pool.
|
||||
for k, v in candidate.dict().items():
|
||||
if k in ENV_ONLY_FIELDS:
|
||||
continue
|
||||
setattr(settings, k, v)
|
||||
cls._current = settings
|
||||
return settings
|
||||
@@ -355,3 +491,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)
|
||||
+553
-86
@@ -1,22 +1,101 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import re
|
||||
import secrets
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncGenerator
|
||||
|
||||
from cashu.core.base import MintQuoteState
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlmodel import col, select
|
||||
from sqlalchemy.orm.attributes import set_committed_value
|
||||
from sqlmodel import col, select, update
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from .core.db import ApiKey, LightningInvoice, create_session, get_session
|
||||
from .core.logging import get_logger
|
||||
from .core.settings import settings
|
||||
from .wallet import get_wallet
|
||||
from .mint import (
|
||||
is_mint_rate_limited,
|
||||
mint_cooldown_remaining,
|
||||
run_mint_operation,
|
||||
)
|
||||
from .wallet import (
|
||||
MintConnectionError,
|
||||
get_wallet,
|
||||
is_mint_connection_error,
|
||||
wallet_operation_guard,
|
||||
)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
lightning_router = APIRouter(prefix="/lightning")
|
||||
|
||||
# Avoid duplicate work within one process. Cross-process settlement is fenced
|
||||
# by claiming a paid quote before minting and by the final conditional update.
|
||||
@dataclass
|
||||
class _InvoiceLockEntry:
|
||||
lock: asyncio.Lock
|
||||
users: int = 0
|
||||
|
||||
|
||||
_invoice_settlement_locks: dict[str, _InvoiceLockEntry] = {}
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _invoice_settlement_lock(invoice_id: str) -> AsyncGenerator[None, None]:
|
||||
"""Serialize one invoice and remove its lock after the last waiter leaves."""
|
||||
|
||||
entry = _invoice_settlement_locks.get(invoice_id)
|
||||
if entry is None:
|
||||
entry = _InvoiceLockEntry(asyncio.Lock())
|
||||
_invoice_settlement_locks[invoice_id] = entry
|
||||
entry.users += 1
|
||||
try:
|
||||
async with entry.lock:
|
||||
yield
|
||||
finally:
|
||||
entry.users -= 1
|
||||
if entry.users == 0 and _invoice_settlement_locks.get(invoice_id) is entry:
|
||||
del _invoice_settlement_locks[invoice_id]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _InvoiceSettlement:
|
||||
id: str
|
||||
payment_hash: str
|
||||
amount_sats: int
|
||||
purpose: str
|
||||
api_key_hash: str | None
|
||||
mint_url: str | None
|
||||
balance_limit: int | None
|
||||
balance_limit_reset: str | None
|
||||
validity_date: int | None
|
||||
|
||||
@classmethod
|
||||
def from_invoice(cls, invoice: LightningInvoice) -> "_InvoiceSettlement":
|
||||
return cls(
|
||||
id=invoice.id,
|
||||
payment_hash=invoice.payment_hash,
|
||||
amount_sats=invoice.amount_sats,
|
||||
purpose=invoice.purpose,
|
||||
api_key_hash=invoice.api_key_hash,
|
||||
mint_url=invoice.mint_url,
|
||||
balance_limit=invoice.balance_limit,
|
||||
balance_limit_reset=invoice.balance_limit_reset,
|
||||
validity_date=invoice.validity_date,
|
||||
)
|
||||
|
||||
|
||||
def _publish_invoice_value(invoice: LightningInvoice, key: str, value: Any) -> None:
|
||||
"""Update a caller view without marking a mapped object dirty."""
|
||||
try:
|
||||
set_committed_value(invoice, key, value)
|
||||
except AttributeError:
|
||||
setattr(invoice, key, value)
|
||||
|
||||
|
||||
class InvoiceCreateRequest(BaseModel):
|
||||
amount_sats: int = Field(gt=0, le=1_000_000, description="Amount in satoshis")
|
||||
@@ -60,16 +139,102 @@ class InvoiceStatusResponse(BaseModel):
|
||||
expires_at: int
|
||||
|
||||
|
||||
_RETRYABLE_INVOICE_STATUSES = ("pending", "settlement_pending")
|
||||
|
||||
|
||||
class InvoiceRecoverRequest(BaseModel):
|
||||
bolt11: str = Field(description="BOLT11 invoice string")
|
||||
|
||||
|
||||
def _trusted_mint_candidates() -> list[str]:
|
||||
return [
|
||||
mint
|
||||
for mint in dict.fromkeys([settings.primary_mint, *settings.cashu_mints])
|
||||
if mint
|
||||
]
|
||||
|
||||
|
||||
async def _request_mint_with_fallback(
|
||||
amount_sats: int,
|
||||
*,
|
||||
allowed_mints: list[str] | None = None,
|
||||
) -> tuple[str, str, str]:
|
||||
"""Request a quote, falling back only among the allowed trusted mints.
|
||||
|
||||
Guards against amount_sats <= 0: the cashu library's PostMintQuoteRequest
|
||||
enforces ``amount > 0`` (Pydantic Field(gt=0)), so passing 0 raises a
|
||||
cryptic validation error deep in the stack. Fail fast with context.
|
||||
"""
|
||||
if amount_sats <= 0:
|
||||
raise ValueError(
|
||||
f"generate_lightning_invoice: amount_sats must be > 0, got {amount_sats}."
|
||||
)
|
||||
tried: list[str] = []
|
||||
trusted = _trusted_mint_candidates()
|
||||
if allowed_mints:
|
||||
# Persisted mint preferences (e.g. an API key's refund_mint_url) must
|
||||
# not outlive the operator's trusted-mint configuration.
|
||||
candidates = [m for m in dict.fromkeys(allowed_mints) if m in trusted]
|
||||
if not candidates:
|
||||
logger.warning(
|
||||
"Requested mints are no longer trusted; falling back to "
|
||||
"configured mints",
|
||||
extra={
|
||||
"requested_mints": list(dict.fromkeys(allowed_mints)),
|
||||
"op_name": "request_mint_invoice",
|
||||
},
|
||||
)
|
||||
candidates = trusted
|
||||
else:
|
||||
candidates = trusted
|
||||
for mint_url in candidates:
|
||||
cooldown = mint_cooldown_remaining(mint_url)
|
||||
if cooldown > 0:
|
||||
tried.append(f"{mint_url}: cooling down")
|
||||
logger.info(
|
||||
"Skipping mint during cooldown",
|
||||
extra={
|
||||
"mint_url": mint_url,
|
||||
"cooldown_seconds": round(cooldown, 2),
|
||||
"op_name": "request_mint_invoice",
|
||||
},
|
||||
)
|
||||
continue
|
||||
try:
|
||||
wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False)
|
||||
quote = await run_mint_operation(
|
||||
lambda: wallet.request_mint(amount_sats),
|
||||
op_name="request_mint_invoice",
|
||||
mint_url=mint_url,
|
||||
retry_on_rate_limit=False,
|
||||
)
|
||||
return quote.request, quote.quote, mint_url
|
||||
except Exception as e:
|
||||
tried.append(f"{mint_url}: {type(e).__name__}")
|
||||
if not is_mint_connection_error(e) and not is_mint_rate_limited(e):
|
||||
raise
|
||||
logger.warning(
|
||||
"request_mint failed, trying fallback mint",
|
||||
extra={
|
||||
"failed_mint": mint_url,
|
||||
"error": str(e),
|
||||
"tried": tried,
|
||||
},
|
||||
)
|
||||
continue
|
||||
raise MintConnectionError(f"All mints failed for request_mint: {tried}")
|
||||
|
||||
|
||||
async def generate_lightning_invoice(
|
||||
amount_sats: int, description: str
|
||||
) -> tuple[str, str]:
|
||||
wallet = await get_wallet(settings.primary_mint, "sat")
|
||||
quote = await wallet.request_mint(amount_sats)
|
||||
return quote.request, quote.quote
|
||||
amount_sats: int,
|
||||
description: str,
|
||||
*,
|
||||
allowed_mints: list[str] | None = None,
|
||||
) -> tuple[str, str, str]:
|
||||
bolt11, payment_hash, mint_url = await _request_mint_with_fallback(
|
||||
amount_sats, allowed_mints=allowed_mints
|
||||
)
|
||||
return bolt11, payment_hash, mint_url
|
||||
|
||||
|
||||
def generate_invoice_id() -> str:
|
||||
@@ -83,6 +248,7 @@ async def create_invoice(
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> InvoiceCreateResponse:
|
||||
api_key_token = _extract_bearer_api_key(authorization) or request.api_key
|
||||
topup_api_key: ApiKey | None = None
|
||||
|
||||
if request.purpose == "topup":
|
||||
if not api_key_token:
|
||||
@@ -93,14 +259,23 @@ async def create_invoice(
|
||||
if not api_key_token.startswith("sk-"):
|
||||
raise HTTPException(status_code=400, detail="Invalid API key format")
|
||||
|
||||
api_key = await session.get(ApiKey, api_key_token[3:])
|
||||
if not api_key:
|
||||
topup_api_key = await session.get(ApiKey, api_key_token[3:])
|
||||
if not topup_api_key:
|
||||
raise HTTPException(status_code=404, detail="API key not found")
|
||||
|
||||
try:
|
||||
description = f"Routstr {request.purpose} {request.amount_sats} sats"
|
||||
bolt11, payment_hash = await generate_lightning_invoice(
|
||||
request.amount_sats, description
|
||||
allowed_mints = None
|
||||
if request.purpose == "topup":
|
||||
assert topup_api_key is not None
|
||||
# A key's liabilities are attributed to a single refund mint. Keep
|
||||
# top-up collateral on that same mint so balances and payouts cannot
|
||||
# misclassify funds held by another mint as owner profit.
|
||||
allowed_mints = [
|
||||
topup_api_key.refund_mint_url or settings.primary_mint
|
||||
]
|
||||
bolt11, payment_hash, mint_url = await generate_lightning_invoice(
|
||||
request.amount_sats, description, allowed_mints=allowed_mints
|
||||
)
|
||||
|
||||
invoice_id = generate_invoice_id()
|
||||
@@ -115,6 +290,7 @@ async def create_invoice(
|
||||
status="pending",
|
||||
api_key_hash=api_key_token[3:] if api_key_token else None,
|
||||
purpose=request.purpose,
|
||||
mint_url=mint_url,
|
||||
balance_limit=request.balance_limit,
|
||||
balance_limit_reset=request.balance_limit_reset,
|
||||
validity_date=request.validity_date,
|
||||
@@ -160,12 +336,12 @@ async def get_invoice_status(
|
||||
if not invoice:
|
||||
raise HTTPException(status_code=404, detail="Invoice not found")
|
||||
|
||||
if invoice.status == "pending":
|
||||
await check_invoice_payment(invoice, session)
|
||||
|
||||
if invoice.status == "pending" and int(time.time()) > invoice.expires_at:
|
||||
invoice.status = "expired"
|
||||
await session.commit()
|
||||
definitively_unpaid = False
|
||||
if invoice.status in _RETRYABLE_INVOICE_STATUSES:
|
||||
definitively_unpaid = await check_invoice_payment(invoice, session)
|
||||
await _expire_invoice_if_authoritatively_unpaid(
|
||||
invoice, session, definitively_unpaid
|
||||
)
|
||||
|
||||
api_key = None
|
||||
if invoice.status == "paid" and invoice.purpose == "create":
|
||||
@@ -199,8 +375,12 @@ async def recover_invoice(
|
||||
if not invoice:
|
||||
raise HTTPException(status_code=404, detail="Invoice not found")
|
||||
|
||||
if invoice.status == "pending":
|
||||
await check_invoice_payment(invoice, session)
|
||||
definitively_unpaid = False
|
||||
if invoice.status in _RETRYABLE_INVOICE_STATUSES:
|
||||
definitively_unpaid = await check_invoice_payment(invoice, session)
|
||||
await _expire_invoice_if_authoritatively_unpaid(
|
||||
invoice, session, definitively_unpaid
|
||||
)
|
||||
|
||||
api_key = None
|
||||
if invoice.status == "paid":
|
||||
@@ -219,114 +399,401 @@ async def recover_invoice(
|
||||
)
|
||||
|
||||
|
||||
async def _claim_paid_invoice_for_settlement(
|
||||
invoice: LightningInvoice,
|
||||
caller_session: AsyncSession,
|
||||
observed_status: str,
|
||||
) -> bool:
|
||||
"""Claim an authoritative paid quote before consuming it at the mint."""
|
||||
if observed_status == "settlement_pending":
|
||||
return True
|
||||
if observed_status != "pending":
|
||||
await _reload_invoice_view(invoice, caller_session)
|
||||
return False
|
||||
|
||||
async with create_session() as claim_session:
|
||||
claim = await claim_session.exec( # type: ignore[call-overload]
|
||||
update(LightningInvoice)
|
||||
.where(
|
||||
col(LightningInvoice.id) == invoice.id,
|
||||
col(LightningInvoice.status) == "pending",
|
||||
)
|
||||
.values(status="settlement_pending")
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
await claim_session.commit()
|
||||
|
||||
if claim.rowcount != 1:
|
||||
await _reload_invoice_view(invoice, caller_session)
|
||||
return False
|
||||
|
||||
_publish_invoice_value(invoice, "status", "settlement_pending")
|
||||
return True
|
||||
|
||||
|
||||
async def check_invoice_payment(
|
||||
invoice: LightningInvoice, session: AsyncSession
|
||||
) -> None:
|
||||
try:
|
||||
wallet = await get_wallet(settings.primary_mint, "sat")
|
||||
|
||||
mint_status = await wallet.get_mint_quote(invoice.payment_hash)
|
||||
|
||||
if mint_status.paid:
|
||||
invoice.status = "paid"
|
||||
invoice.paid_at = int(time.time())
|
||||
|
||||
if invoice.purpose == "create":
|
||||
api_key = await create_api_key_from_invoice(invoice, session)
|
||||
invoice.api_key_hash = api_key.hashed_key
|
||||
elif invoice.purpose == "topup" and invoice.api_key_hash:
|
||||
await topup_api_key_from_invoice(invoice, session)
|
||||
) -> bool:
|
||||
"""Settle an invoice and report whether its quote is definitively unpaid.
|
||||
|
||||
False covers paid, pending, and ambiguous transport/DB outcomes so callers
|
||||
never expire a quote merely because reconciliation could not complete.
|
||||
"""
|
||||
async with _invoice_settlement_lock(invoice.id), wallet_operation_guard():
|
||||
minted = False
|
||||
payment_confirmed = False
|
||||
try:
|
||||
# Snapshot the row and end the caller's read transaction before any
|
||||
# potentially slow mint I/O. All final DB mutations use owned,
|
||||
# short-lived sessions below.
|
||||
await session.refresh(invoice)
|
||||
if invoice.status not in _RETRYABLE_INVOICE_STATUSES:
|
||||
await session.commit()
|
||||
return False
|
||||
observed_status = invoice.status
|
||||
settlement = _InvoiceSettlement.from_invoice(invoice)
|
||||
await session.commit()
|
||||
|
||||
mint_url = settlement.mint_url or settings.primary_mint
|
||||
wallet = await get_wallet(mint_url, "sat")
|
||||
mint_status = await run_mint_operation(
|
||||
lambda: wallet.get_mint_quote(settlement.payment_hash),
|
||||
op_name="get_mint_quote",
|
||||
mint_url=mint_url,
|
||||
)
|
||||
if not mint_status.paid:
|
||||
return getattr(mint_status, "state", None) == MintQuoteState.unpaid
|
||||
payment_confirmed = True
|
||||
|
||||
# Fence expiry and other workers before consuming the paid quote.
|
||||
# If a concurrent expiry/finalization won, this worker must not mint.
|
||||
if not await _claim_paid_invoice_for_settlement(
|
||||
invoice, session, observed_status
|
||||
):
|
||||
return False
|
||||
|
||||
# Reject a paid top-up whose target was pruned before redeeming its
|
||||
# single-use quote. The validation session is closed before mint I/O.
|
||||
if settlement.purpose == "topup":
|
||||
if not settlement.api_key_hash:
|
||||
raise ValueError("No API key associated with topup invoice")
|
||||
async with create_session() as validation_session:
|
||||
target = await validation_session.get(
|
||||
ApiKey, settlement.api_key_hash
|
||||
)
|
||||
if target is None:
|
||||
terminal = await validation_session.exec( # type: ignore[call-overload]
|
||||
update(LightningInvoice)
|
||||
.where(
|
||||
col(LightningInvoice.id) == settlement.id,
|
||||
col(LightningInvoice.status).in_(
|
||||
_RETRYABLE_INVOICE_STATUSES
|
||||
),
|
||||
)
|
||||
.values(status="reconciliation_required")
|
||||
)
|
||||
await validation_session.commit()
|
||||
if terminal.rowcount == 1:
|
||||
_publish_invoice_value(
|
||||
invoice, "status", "reconciliation_required"
|
||||
)
|
||||
else:
|
||||
await _reload_invoice_view(invoice, session)
|
||||
logger.critical(
|
||||
"Paid topup invoice target API key was not found; reconciliation required",
|
||||
extra={"invoice_id": settlement.id},
|
||||
)
|
||||
return False
|
||||
|
||||
# Quote-linked proof verification makes an ambiguous mint response
|
||||
# retryable without crediting unrelated wallet balance growth.
|
||||
await _mint_invoice_quote(wallet, settlement)
|
||||
minted = True
|
||||
|
||||
paid_at = int(time.time())
|
||||
async with create_session() as finalization_session:
|
||||
settled, api_key_hash = await _finalize_invoice_settlement(
|
||||
settlement, finalization_session, paid_at
|
||||
)
|
||||
if not settled:
|
||||
await _reload_invoice_view(invoice, session)
|
||||
return False
|
||||
|
||||
_publish_invoice_value(invoice, "status", "paid")
|
||||
_publish_invoice_value(invoice, "paid_at", paid_at)
|
||||
_publish_invoice_value(invoice, "api_key_hash", api_key_hash)
|
||||
logger.info(
|
||||
"Lightning invoice paid",
|
||||
extra={
|
||||
"invoice_id": invoice.id,
|
||||
"amount_sats": invoice.amount_sats,
|
||||
"purpose": invoice.purpose,
|
||||
"api_key_hash": invoice.api_key_hash[:8] + "..."
|
||||
if invoice.api_key_hash
|
||||
"invoice_id": settlement.id,
|
||||
"amount_sats": settlement.amount_sats,
|
||||
"purpose": settlement.purpose,
|
||||
"api_key_hash": api_key_hash[:8] + "..."
|
||||
if api_key_hash
|
||||
else None,
|
||||
},
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to check invoice payment: {e}")
|
||||
return False
|
||||
except BaseException as error:
|
||||
# Never roll back the caller-owned session: doing so expires invoice
|
||||
# and sibling ORM objects. Owned sessions roll themselves back.
|
||||
if payment_confirmed and invoice.status != "settlement_pending":
|
||||
try:
|
||||
async with create_session() as state_session:
|
||||
pending = await state_session.exec( # type: ignore[call-overload]
|
||||
update(LightningInvoice)
|
||||
.where(
|
||||
col(LightningInvoice.id) == invoice.id,
|
||||
col(LightningInvoice.status).in_(
|
||||
_RETRYABLE_INVOICE_STATUSES
|
||||
),
|
||||
)
|
||||
.values(status="settlement_pending")
|
||||
)
|
||||
await state_session.commit()
|
||||
if pending.rowcount == 1:
|
||||
_publish_invoice_value(
|
||||
invoice, "status", "settlement_pending"
|
||||
)
|
||||
except Exception as state_error:
|
||||
logger.critical(
|
||||
"Paid invoice reconciliation state could not be persisted",
|
||||
extra={"invoice_id": invoice.id, "error": str(state_error)},
|
||||
)
|
||||
if minted:
|
||||
logger.critical(
|
||||
"Invoice mint succeeded but DB finalization failed; reconciliation required",
|
||||
extra={"invoice_id": invoice.id, "purpose": invoice.purpose},
|
||||
)
|
||||
try:
|
||||
await _reload_invoice_view(invoice, session)
|
||||
except Exception:
|
||||
pass
|
||||
if not isinstance(error, Exception):
|
||||
raise
|
||||
logger.error(f"Failed to check invoice payment: {error}")
|
||||
return False
|
||||
|
||||
|
||||
async def create_api_key_from_invoice(
|
||||
invoice: LightningInvoice, session: AsyncSession
|
||||
) -> ApiKey:
|
||||
wallet = await get_wallet(settings.primary_mint, "sat")
|
||||
await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash)
|
||||
def _is_outputs_already_signed(error: BaseException) -> bool:
|
||||
message = str(error)
|
||||
return bool(
|
||||
re.search(
|
||||
r"\boutputs?\s+(?:have\s+)?already\s+(?:been\s+)?signed(?:\s+before)?\b",
|
||||
message,
|
||||
re.IGNORECASE,
|
||||
)
|
||||
and re.search(r"\bcode\s*:\s*11003\b", message, re.IGNORECASE)
|
||||
)
|
||||
|
||||
|
||||
def _invoice_quote_proof_amount(wallet: Any, quote_id: str) -> int:
|
||||
"""Return spendable wallet value minted by one Lightning quote."""
|
||||
return sum(
|
||||
proof.amount
|
||||
for proof in wallet.proofs
|
||||
if proof.mint_id == quote_id and not proof.reserved
|
||||
)
|
||||
|
||||
|
||||
async def _mint_invoice_quote(
|
||||
wallet: Any, invoice: LightningInvoice | _InvoiceSettlement
|
||||
) -> None:
|
||||
"""Mint a paid quote, proving quote-linked outputs before DB credit."""
|
||||
mint_url = invoice.mint_url or settings.primary_mint
|
||||
await wallet.load_proofs(reload=True)
|
||||
if _invoice_quote_proof_amount(wallet, invoice.payment_hash) >= invoice.amount_sats:
|
||||
return
|
||||
|
||||
try:
|
||||
await run_mint_operation(
|
||||
lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash),
|
||||
op_name=f"invoice_mint_{invoice.purpose}",
|
||||
mint_url=mint_url,
|
||||
retry_timeouts=False,
|
||||
)
|
||||
except Exception as error:
|
||||
if not _is_outputs_already_signed(error):
|
||||
raise
|
||||
|
||||
for keyset_id in wallet.keysets:
|
||||
await wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25)
|
||||
await wallet.load_proofs(reload=True)
|
||||
recovered = _invoice_quote_proof_amount(wallet, invoice.payment_hash)
|
||||
if recovered < invoice.amount_sats:
|
||||
raise RuntimeError(
|
||||
"Invoice outputs were already signed but quote-linked recovery returned "
|
||||
f"{recovered} sats; expected at least {invoice.amount_sats}"
|
||||
) from error
|
||||
else:
|
||||
await wallet.load_proofs(reload=True)
|
||||
minted_amount = _invoice_quote_proof_amount(wallet, invoice.payment_hash)
|
||||
if minted_amount < invoice.amount_sats:
|
||||
raise RuntimeError(
|
||||
"Invoice mint succeeded but quote-linked proofs total "
|
||||
f"{minted_amount} sats; expected at least {invoice.amount_sats}"
|
||||
)
|
||||
|
||||
|
||||
def _invoice_api_key_hash(invoice: LightningInvoice | _InvoiceSettlement) -> str:
|
||||
dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}"
|
||||
hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest()
|
||||
return hashlib.sha256(dummy_token.encode()).hexdigest()
|
||||
|
||||
|
||||
async def _create_api_key_record(
|
||||
invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession
|
||||
) -> ApiKey:
|
||||
mint_url = invoice.mint_url or settings.primary_mint
|
||||
api_key = ApiKey(
|
||||
hashed_key=hashed_key,
|
||||
balance=invoice.amount_sats * 1000, # Convert to msats
|
||||
hashed_key=_invoice_api_key_hash(invoice),
|
||||
balance=invoice.amount_sats * 1000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url=settings.primary_mint,
|
||||
refund_mint_url=mint_url,
|
||||
balance_limit=invoice.balance_limit,
|
||||
balance_limit_reset=invoice.balance_limit_reset,
|
||||
validity_date=invoice.validity_date,
|
||||
)
|
||||
|
||||
session.add(api_key)
|
||||
await session.flush()
|
||||
|
||||
return api_key
|
||||
|
||||
|
||||
async def topup_api_key_from_invoice(
|
||||
invoice: LightningInvoice, session: AsyncSession
|
||||
async def _topup_api_key_record(
|
||||
invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession
|
||||
) -> None:
|
||||
wallet = await get_wallet(settings.primary_mint, "sat")
|
||||
await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash)
|
||||
|
||||
if not invoice.api_key_hash:
|
||||
raise ValueError("No API key associated with topup invoice")
|
||||
|
||||
api_key = await session.get(ApiKey, invoice.api_key_hash)
|
||||
if not api_key:
|
||||
result = await session.exec( # type: ignore[call-overload]
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == invoice.api_key_hash)
|
||||
.values(balance=col(ApiKey.balance) + invoice.amount_sats * 1000)
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
if result.rowcount != 1:
|
||||
raise ValueError("Associated API key not found")
|
||||
|
||||
api_key.balance += invoice.amount_sats * 1000 # Convert to msats
|
||||
await session.flush()
|
||||
|
||||
async def _finalize_invoice_settlement(
|
||||
invoice: _InvoiceSettlement, session: AsyncSession, paid_at: int
|
||||
) -> tuple[bool, str | None]:
|
||||
"""Atomically fence and apply one invoice credit in the provided owned session."""
|
||||
api_key_hash = (
|
||||
_invoice_api_key_hash(invoice)
|
||||
if invoice.purpose == "create"
|
||||
else invoice.api_key_hash
|
||||
)
|
||||
claim = await session.exec( # type: ignore[call-overload]
|
||||
update(LightningInvoice)
|
||||
.where(col(LightningInvoice.id) == invoice.id)
|
||||
.where(
|
||||
col(LightningInvoice.status).in_(_RETRYABLE_INVOICE_STATUSES)
|
||||
)
|
||||
.values(status="paid", paid_at=paid_at, api_key_hash=api_key_hash)
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
if claim.rowcount != 1:
|
||||
await session.rollback()
|
||||
return False, None
|
||||
|
||||
if invoice.purpose == "create":
|
||||
await _create_api_key_record(invoice, session)
|
||||
elif invoice.purpose == "topup":
|
||||
await _topup_api_key_record(invoice, session)
|
||||
else:
|
||||
raise ValueError(f"Unsupported invoice purpose: {invoice.purpose}")
|
||||
await session.commit()
|
||||
return True, api_key_hash
|
||||
|
||||
|
||||
INVOICE_WATCH_INTERVAL_SECONDS = 5
|
||||
async def _reload_invoice_view(
|
||||
invoice: LightningInvoice, _caller_session: AsyncSession
|
||||
) -> None:
|
||||
"""Publish committed invoice state without touching the caller transaction."""
|
||||
async with create_session() as reload_session:
|
||||
stored = await reload_session.get(LightningInvoice, invoice.id)
|
||||
if stored is None:
|
||||
return
|
||||
status = stored.status
|
||||
paid_at = stored.paid_at
|
||||
api_key_hash = stored.api_key_hash
|
||||
await reload_session.commit()
|
||||
_publish_invoice_value(invoice, "status", status)
|
||||
_publish_invoice_value(invoice, "paid_at", paid_at)
|
||||
_publish_invoice_value(invoice, "api_key_hash", api_key_hash)
|
||||
|
||||
|
||||
async def _expire_invoice_if_authoritatively_unpaid(
|
||||
invoice: LightningInvoice,
|
||||
caller_session: AsyncSession,
|
||||
definitively_unpaid: bool,
|
||||
) -> bool:
|
||||
"""Expire one overdue unpaid invoice without overwriting concurrent settlement."""
|
||||
if (
|
||||
not definitively_unpaid
|
||||
or invoice.status != "pending"
|
||||
or int(time.time()) <= invoice.expires_at
|
||||
):
|
||||
return False
|
||||
|
||||
async with create_session() as expiry_session:
|
||||
expired = await expiry_session.exec( # type: ignore[call-overload]
|
||||
update(LightningInvoice)
|
||||
.where(
|
||||
col(LightningInvoice.id) == invoice.id,
|
||||
col(LightningInvoice.status) == "pending",
|
||||
)
|
||||
.values(status="expired")
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
await expiry_session.commit()
|
||||
|
||||
if expired.rowcount == 1:
|
||||
_publish_invoice_value(invoice, "status", "expired")
|
||||
return True
|
||||
|
||||
await _reload_invoice_view(invoice, caller_session)
|
||||
return False
|
||||
|
||||
|
||||
async def _credit_topup_record(
|
||||
invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession
|
||||
) -> None:
|
||||
await _topup_api_key_record(invoice, session)
|
||||
|
||||
|
||||
# Nutshell mints throttle Lightning backend lookups to once per 10s per
|
||||
# quote, so polling faster just burns the global request budget for nothing.
|
||||
INVOICE_WATCH_INTERVAL_SECONDS = 10
|
||||
INVOICE_WATCH_BATCH_LIMIT = 100
|
||||
|
||||
|
||||
async def periodic_invoice_watcher() -> None:
|
||||
"""Background task: detect paid Lightning invoices and credit balances.
|
||||
async def _process_invoice_watch_batch(session: AsyncSession) -> None:
|
||||
result = await session.exec(
|
||||
select(LightningInvoice)
|
||||
.where(
|
||||
col(LightningInvoice.status).in_(_RETRYABLE_INVOICE_STATUSES)
|
||||
)
|
||||
.limit(INVOICE_WATCH_BATCH_LIMIT)
|
||||
)
|
||||
for invoice in result.all():
|
||||
try:
|
||||
definitively_unpaid = await check_invoice_payment(invoice, session)
|
||||
await _expire_invoice_if_authoritatively_unpaid(
|
||||
invoice, session, definitively_unpaid
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Invoice watcher failed for invoice",
|
||||
extra={"invoice_id": invoice.id, "error": str(e)},
|
||||
)
|
||||
|
||||
Removes the need for clients to poll the status endpoint after paying.
|
||||
"""
|
||||
|
||||
async def periodic_invoice_watcher() -> None:
|
||||
"""Background task: detect paid Lightning invoices and credit balances."""
|
||||
while True:
|
||||
try:
|
||||
async with create_session() as session:
|
||||
now = int(time.time())
|
||||
result = await session.exec(
|
||||
select(LightningInvoice)
|
||||
.where(
|
||||
LightningInvoice.status == "pending",
|
||||
col(LightningInvoice.expires_at) > now,
|
||||
)
|
||||
.limit(INVOICE_WATCH_BATCH_LIMIT)
|
||||
)
|
||||
pending = result.all()
|
||||
for invoice in pending:
|
||||
try:
|
||||
await check_invoice_payment(invoice, session)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Invoice watcher failed for invoice",
|
||||
extra={"invoice_id": invoice.id, "error": str(e)},
|
||||
)
|
||||
await _process_invoice_watch_batch(session)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
||||
+343
@@ -0,0 +1,343 @@
|
||||
"""Shared policy for bounded, rate-aware Cashu mint API operations."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import socket
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from contextvars import ContextVar
|
||||
from typing import Any, AsyncGenerator, Awaitable, Callable
|
||||
|
||||
import httpx
|
||||
|
||||
from .core.logging import get_logger
|
||||
from .core.settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
MINT_TRANSPORT_EXCEPTIONS: tuple[type[BaseException], ...] = (
|
||||
httpx.NetworkError,
|
||||
httpx.TimeoutException,
|
||||
ConnectionError,
|
||||
socket.gaierror,
|
||||
asyncio.TimeoutError,
|
||||
)
|
||||
|
||||
MINT_TRANSPORT_COOLDOWN_SECONDS = 30.0
|
||||
_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS = 60.0
|
||||
_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS = 7 * 60 * 60
|
||||
|
||||
_fail_fast_depth: ContextVar[int] = ContextVar("mint_fail_fast_depth", default=0)
|
||||
|
||||
|
||||
class MintRateLimitedError(httpx.HTTPStatusError):
|
||||
"""Typed boundary error preserving a Cashu mint's HTTP 429 response."""
|
||||
|
||||
|
||||
class MintCooldownError(Exception):
|
||||
"""A mint is cooling down and this operation must not wait."""
|
||||
|
||||
def __init__(self, mint_url: str, retry_after_seconds: float):
|
||||
self.mint_url = mint_url
|
||||
self.retry_after_seconds = max(0.0, retry_after_seconds)
|
||||
super().__init__(
|
||||
f"Mint {mint_url} is cooling down; retry after "
|
||||
f"{self.retry_after_seconds:.2f}s"
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def fail_fast_mint_operations() -> AsyncGenerator[None, None]:
|
||||
"""Make mint cooldown/probe waits fail fast in the current task.
|
||||
|
||||
Wallet mutation code holds a process-wide file lock. It enters this scope so
|
||||
an existing mint cooldown can never turn that lock into a multi-hour wait.
|
||||
"""
|
||||
|
||||
token = _fail_fast_depth.set(_fail_fast_depth.get() + 1)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_fail_fast_depth.reset(token)
|
||||
|
||||
|
||||
class MintRateGuard:
|
||||
"""Limit concurrency and remember per-mint cooldown/probe state."""
|
||||
|
||||
_guards: dict[str, "MintRateGuard"] = {}
|
||||
|
||||
@classmethod
|
||||
def get(cls, mint_url: str) -> "MintRateGuard":
|
||||
concurrency = settings.mint_max_concurrency
|
||||
guard = cls._guards.get(mint_url)
|
||||
if guard is None or guard._max_concurrency != concurrency:
|
||||
previous = guard
|
||||
guard = cls(mint_url, concurrency)
|
||||
if previous is not None:
|
||||
# Concurrency changed at runtime: keep the live cooldown/backoff
|
||||
# state so an active 429 cooldown is not silently discarded.
|
||||
guard._cooldown_until = previous._cooldown_until
|
||||
guard._cooldown_reason = previous._cooldown_reason
|
||||
guard._consecutive_rate_limits = previous._consecutive_rate_limits
|
||||
guard._needs_probe = previous._needs_probe
|
||||
cls._guards[mint_url] = guard
|
||||
return guard
|
||||
|
||||
def __init__(self, mint_url: str, max_concurrency: int):
|
||||
self._mint_url = mint_url
|
||||
self._max_concurrency = max_concurrency
|
||||
self._semaphore = (
|
||||
asyncio.Semaphore(max_concurrency) if max_concurrency > 0 else None
|
||||
)
|
||||
self._cooldown_until = 0.0
|
||||
self._cooldown_reason: str | None = None
|
||||
self._consecutive_rate_limits = 0
|
||||
self._needs_probe = False
|
||||
self._probe_lock = asyncio.Lock()
|
||||
|
||||
def apply_cooldown(self, delay: float, *, reason: str | None = None) -> None:
|
||||
deadline = time.monotonic() + max(0.0, delay)
|
||||
if deadline >= self._cooldown_until:
|
||||
self._cooldown_until = deadline
|
||||
if reason is not None:
|
||||
self._cooldown_reason = reason
|
||||
elif self._cooldown_reason is None and reason is not None:
|
||||
self._cooldown_reason = reason
|
||||
self._needs_probe = True
|
||||
|
||||
def apply_rate_limit_cooldown(self, retry_after: float | None = None) -> float:
|
||||
remaining = self.cooldown_remaining()
|
||||
if remaining > 0 and self._cooldown_reason == "rate_limited":
|
||||
minimum = min(
|
||||
_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS,
|
||||
max(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, retry_after or 0.0),
|
||||
)
|
||||
if minimum > remaining:
|
||||
self.apply_cooldown(minimum, reason="rate_limited")
|
||||
return minimum
|
||||
return remaining
|
||||
|
||||
self._consecutive_rate_limits += 1
|
||||
base = max(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, retry_after or 0.0)
|
||||
multiplier = 2 ** min(self._consecutive_rate_limits - 1, 10)
|
||||
delay = min(_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS, base * multiplier)
|
||||
self.apply_cooldown(delay, reason="rate_limited")
|
||||
return delay
|
||||
|
||||
def cooldown_remaining(self) -> float:
|
||||
return max(0.0, self._cooldown_until - time.monotonic())
|
||||
|
||||
def cooldown_reason(self) -> str | None:
|
||||
return self._cooldown_reason if self.cooldown_remaining() > 0 else None
|
||||
|
||||
def _raise_if_wait_forbidden(self) -> None:
|
||||
remaining = self.cooldown_remaining()
|
||||
if _fail_fast_depth.get() and remaining > 0:
|
||||
raise MintCooldownError(self._mint_url, remaining)
|
||||
|
||||
async def _wait_for_cooldown(self) -> None:
|
||||
while True:
|
||||
self._raise_if_wait_forbidden()
|
||||
deadline = self._cooldown_until
|
||||
wait = max(0.0, deadline - time.monotonic())
|
||||
if wait <= 0:
|
||||
return
|
||||
logger.debug(
|
||||
"Mint rate guard: cooling down",
|
||||
extra={"mint_url": self._mint_url, "wait_seconds": round(wait, 2)},
|
||||
)
|
||||
await asyncio.sleep(wait)
|
||||
if self._cooldown_until <= deadline:
|
||||
return
|
||||
|
||||
async def _run_probe(self, factory: Callable[[], Awaitable[Any]]) -> Any:
|
||||
await self._wait_for_cooldown()
|
||||
logger.info(
|
||||
"Mint cooldown ended; sending one probe request",
|
||||
extra={"event": "mint_cooldown_probe_started", "mint_url": self._mint_url},
|
||||
)
|
||||
try:
|
||||
result = await factory()
|
||||
except Exception as error:
|
||||
if is_mint_rate_limited(error):
|
||||
retry_after = None
|
||||
if isinstance(error, httpx.HTTPStatusError):
|
||||
retry_after = parse_retry_after(error.response.headers)
|
||||
self.apply_rate_limit_cooldown(retry_after)
|
||||
else:
|
||||
self.apply_cooldown(1.0)
|
||||
logger.warning(
|
||||
"Mint cooldown probe failed",
|
||||
extra={
|
||||
"event": "mint_cooldown_probe_failed",
|
||||
"mint_url": self._mint_url,
|
||||
"error": str(error),
|
||||
"error_type": type(error).__name__,
|
||||
"cooldown_seconds": round(self.cooldown_remaining(), 2),
|
||||
"consecutive_rate_limits": self._consecutive_rate_limits,
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
self._needs_probe = False
|
||||
self._cooldown_until = 0.0
|
||||
self._cooldown_reason = None
|
||||
self._consecutive_rate_limits = 0
|
||||
logger.info(
|
||||
"Mint cooldown probe succeeded; restoring normal concurrency",
|
||||
extra={
|
||||
"event": "mint_cooldown_probe_succeeded",
|
||||
"mint_url": self._mint_url,
|
||||
},
|
||||
)
|
||||
return result
|
||||
|
||||
async def run(self, factory: Callable[[], Awaitable[Any]]) -> Any:
|
||||
while True:
|
||||
self._raise_if_wait_forbidden()
|
||||
if self._needs_probe or self.cooldown_remaining() > 0:
|
||||
if _fail_fast_depth.get() and self._probe_lock.locked():
|
||||
raise MintCooldownError(self._mint_url, self.cooldown_remaining())
|
||||
async with self._probe_lock:
|
||||
self._raise_if_wait_forbidden()
|
||||
if self.cooldown_remaining() > 0:
|
||||
self._needs_probe = True
|
||||
if self._needs_probe:
|
||||
return await self._run_probe(factory)
|
||||
continue
|
||||
|
||||
if self._semaphore is None:
|
||||
return await factory()
|
||||
async with self._semaphore:
|
||||
self._raise_if_wait_forbidden()
|
||||
if self._needs_probe:
|
||||
continue
|
||||
return await factory()
|
||||
|
||||
|
||||
def mint_cooldown_remaining(mint_url: str) -> float:
|
||||
return MintRateGuard.get(mint_url).cooldown_remaining()
|
||||
|
||||
|
||||
def mint_cooldown_reason(mint_url: str) -> str | None:
|
||||
return MintRateGuard.get(mint_url).cooldown_reason()
|
||||
|
||||
|
||||
def is_mint_rate_limited(error: BaseException) -> bool:
|
||||
"""Return whether an exception chain represents HTTP 429/cooldown."""
|
||||
|
||||
current: BaseException | None = error
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
if isinstance(current, MintCooldownError):
|
||||
return True
|
||||
if isinstance(current, httpx.HTTPStatusError):
|
||||
if current.response.status_code == 429:
|
||||
return True
|
||||
current = current.__cause__ or current.__context__
|
||||
return False
|
||||
|
||||
|
||||
def parse_retry_after(headers: Any) -> float | None:
|
||||
raw = headers.get("retry-after") or headers.get("Retry-After")
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
return float(str(raw).strip())
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
async def run_mint_operation(
|
||||
factory: Callable[[], Awaitable[Any]],
|
||||
*,
|
||||
op_name: str = "mint_operation",
|
||||
mint_url: str = "",
|
||||
retry_timeouts: bool = True,
|
||||
retry_on_rate_limit: bool = True,
|
||||
) -> Any:
|
||||
"""Run one mint operation with bounded concurrency and adaptive cooldown."""
|
||||
|
||||
guard = MintRateGuard.get(mint_url) if mint_url else None
|
||||
timeout = settings.mint_operation_timeout_seconds
|
||||
max_attempts = settings.mint_retry_max_attempts + 1
|
||||
|
||||
async def timed_factory() -> Any:
|
||||
if timeout > 0:
|
||||
return await asyncio.wait_for(factory(), timeout=timeout)
|
||||
return await factory()
|
||||
|
||||
async def invoke() -> Any:
|
||||
if guard is not None:
|
||||
return await guard.run(timed_factory)
|
||||
return await timed_factory()
|
||||
|
||||
for attempt in range(max_attempts):
|
||||
try:
|
||||
return await invoke()
|
||||
except MintCooldownError:
|
||||
raise
|
||||
except (asyncio.TimeoutError, httpx.TimeoutException) as exc:
|
||||
if retry_timeouts and attempt < max_attempts - 1:
|
||||
backoff = (2**attempt) + (time.monotonic() % 1.0)
|
||||
logger.warning(
|
||||
"Mint operation timed out, retrying",
|
||||
extra={
|
||||
"op_name": op_name,
|
||||
"mint_url": mint_url,
|
||||
"attempt": attempt + 1,
|
||||
"backoff_seconds": round(backoff, 2),
|
||||
},
|
||||
)
|
||||
await asyncio.sleep(backoff)
|
||||
continue
|
||||
raise httpx.TimeoutException(
|
||||
f"{op_name} timed out (attempts: {attempt + 1})"
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
if not is_mint_rate_limited(exc):
|
||||
raise
|
||||
|
||||
backoff = (2**attempt) + (time.monotonic() % 1.0)
|
||||
if isinstance(exc, httpx.HTTPStatusError):
|
||||
retry_after = parse_retry_after(exc.response.headers)
|
||||
if retry_after is not None:
|
||||
backoff = max(retry_after, backoff)
|
||||
cooldown = backoff
|
||||
if guard is not None:
|
||||
cooldown = guard.apply_rate_limit_cooldown(backoff)
|
||||
|
||||
if not retry_on_rate_limit:
|
||||
logger.warning(
|
||||
"Mint rate-limited, skipping retries for fallback",
|
||||
extra={
|
||||
"op_name": op_name,
|
||||
"mint_url": mint_url,
|
||||
"cooldown_seconds": round(cooldown, 2),
|
||||
"consecutive_rate_limits": guard._consecutive_rate_limits
|
||||
if guard is not None
|
||||
else attempt + 1,
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
if attempt >= max_attempts - 1:
|
||||
raise
|
||||
logger.warning(
|
||||
"Mint rate-limited, applying cooldown",
|
||||
extra={
|
||||
"op_name": op_name,
|
||||
"mint_url": mint_url,
|
||||
"attempt": attempt + 1,
|
||||
"cooldown_seconds": round(cooldown, 2),
|
||||
"consecutive_rate_limits": guard._consecutive_rate_limits
|
||||
if guard is not None
|
||||
else attempt + 1,
|
||||
},
|
||||
)
|
||||
if guard is None:
|
||||
await asyncio.sleep(cooldown)
|
||||
|
||||
raise RuntimeError(f"{op_name}: exhausted retries unexpectedly")
|
||||
@@ -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,12 +172,39 @@ 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")
|
||||
)
|
||||
cache_pricing_rates: tuple[float, float, float, float] | None = None
|
||||
if cache_read_tokens > 0 or cache_creation_tokens > 0:
|
||||
try:
|
||||
cache_pricing_rates = _get_pricing_rates(
|
||||
response_data, model_obj, provider_fee
|
||||
)
|
||||
except ValueError:
|
||||
logger.warning(
|
||||
"Cache pricing unavailable for USD cost breakdown; "
|
||||
"leaving cache cost components unknown",
|
||||
extra={"model": response_data.get("model", "unknown")},
|
||||
)
|
||||
if cache_pricing_rates is None and settings.fixed_pricing:
|
||||
fixed_input_rate = (
|
||||
float(settings.fixed_per_1k_input_tokens) * 1000.0
|
||||
)
|
||||
cache_pricing_rates = (
|
||||
fixed_input_rate,
|
||||
float(settings.fixed_per_1k_output_tokens) * 1000.0,
|
||||
fixed_input_rate,
|
||||
fixed_input_rate,
|
||||
)
|
||||
return _calculate_from_usd_cost(
|
||||
usd_cost,
|
||||
input_usd,
|
||||
@@ -172,6 +214,8 @@ async def calculate_cost(
|
||||
cache_creation_tokens,
|
||||
output_tokens,
|
||||
response_data,
|
||||
provider_fee,
|
||||
cache_pricing_rates,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
@@ -185,7 +229,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 +303,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 +321,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 +346,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,19 +473,31 @@ def _calculate_from_usd_cost(
|
||||
cache_creation_tokens: int,
|
||||
output_tokens: int,
|
||||
response_data: dict,
|
||||
provider_fee: float | None,
|
||||
pricing_rates: tuple[float, float, float, float] | None = 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
|
||||
sats_per_usd = 1.0 / sats_usd_price()
|
||||
cost_in_sats = usd_cost * sats_per_usd
|
||||
cost_in_msats = math.ceil(cost_in_sats * 1000)
|
||||
raw_cost_msats = cost_in_sats * 1000
|
||||
cost_in_msats = math.ceil(raw_cost_msats)
|
||||
raw_input_msats = 0.0
|
||||
|
||||
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
|
||||
# Match the token-priced path: truncate the visible output component
|
||||
# and assign the authoritative total's rounding remainder to input.
|
||||
output_msats = math.floor(cost_in_msats * output_usd / component_usd)
|
||||
input_msats = cost_in_msats - output_msats
|
||||
raw_input_msats = raw_cost_msats * input_usd / component_usd
|
||||
else:
|
||||
effective_input_tokens = (
|
||||
input_tokens + cache_read_tokens + cache_creation_tokens
|
||||
@@ -380,6 +509,38 @@ def _calculate_from_usd_cost(
|
||||
else 0
|
||||
)
|
||||
output_msats = cost_in_msats - input_msats
|
||||
raw_input_msats = (
|
||||
raw_cost_msats * effective_input_tokens / total_tokens
|
||||
if total_tokens > 0
|
||||
else 0.0
|
||||
)
|
||||
|
||||
# Preserve the same cache-rate ratios as the token-priced path while the
|
||||
# upstream USD total remains authoritative. Cache values are informational
|
||||
# subcomponents of the inclusive input cost.
|
||||
cache_read_msats = 0
|
||||
cache_creation_msats = 0
|
||||
if pricing_rates is not None:
|
||||
input_rate, _, cache_read_rate, cache_creation_rate = pricing_rates
|
||||
regular_weight = input_tokens * input_rate
|
||||
cache_read_weight = cache_read_tokens * cache_read_rate
|
||||
cache_creation_weight = cache_creation_tokens * cache_creation_rate
|
||||
total_input_weight = (
|
||||
regular_weight + cache_read_weight + cache_creation_weight
|
||||
)
|
||||
if total_input_weight > 0:
|
||||
cache_read_msats = int(
|
||||
round(
|
||||
raw_input_msats * cache_read_weight / total_input_weight,
|
||||
3,
|
||||
)
|
||||
)
|
||||
cache_creation_msats = int(
|
||||
round(
|
||||
raw_input_msats * cache_creation_weight / total_input_weight,
|
||||
3,
|
||||
)
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Using cost from usage data/details",
|
||||
@@ -387,6 +548,8 @@ def _calculate_from_usd_cost(
|
||||
"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 +564,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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -18,7 +18,6 @@ from ..wallet import deserialize_token_from_string
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None:
|
||||
if x_cashu := headers.get("x-cashu", None):
|
||||
cashu_token = x_cashu
|
||||
@@ -243,7 +242,7 @@ async def calculate_discounted_max_cost(
|
||||
},
|
||||
)
|
||||
|
||||
return max(0, adjusted)
|
||||
return max(settings.min_request_msat, adjusted)
|
||||
|
||||
|
||||
def estimate_tokens(messages: list) -> int:
|
||||
|
||||
@@ -4,8 +4,15 @@ import math
|
||||
from typing import TypedDict
|
||||
|
||||
import httpx
|
||||
from cashu.core.base import MeltQuoteState
|
||||
from cashu.wallet.wallet import Proof, Wallet
|
||||
|
||||
from ..mint import (
|
||||
MINT_TRANSPORT_EXCEPTIONS,
|
||||
is_mint_rate_limited,
|
||||
run_mint_operation,
|
||||
)
|
||||
|
||||
try:
|
||||
from bech32 import bech32_decode, convertbits # type: ignore
|
||||
except ModuleNotFoundError: # pragma: no cover – allow runtime miss
|
||||
@@ -25,6 +32,15 @@ class LNURLError(Exception):
|
||||
"""LNURL related errors."""
|
||||
|
||||
|
||||
class MeltOutcomeAmbiguousError(LNURLError):
|
||||
"""A melt was dispatched but its final outcome could not be confirmed.
|
||||
|
||||
Callers must NOT treat this as a clean failure: the payment may still
|
||||
settle, so debits backing it must be kept until reconciliation confirms
|
||||
the true outcome.
|
||||
"""
|
||||
|
||||
|
||||
async def decode_lnurl(lnurl: str) -> str:
|
||||
"""Decode LNURL to get the actual URL.
|
||||
|
||||
@@ -215,15 +231,62 @@ async def raw_send_to_lnurl(
|
||||
lnurl_data["callback_url"], final_amount
|
||||
)
|
||||
|
||||
melt_quote_resp = await wallet.melt_quote(invoice=bolt11_invoice)
|
||||
melt_quote_resp = await run_mint_operation(
|
||||
lambda: wallet.melt_quote(invoice=bolt11_invoice),
|
||||
op_name="lnurl_melt_quote",
|
||||
mint_url=str(wallet.url),
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
return final_amount
|
||||
try:
|
||||
melt_response = await run_mint_operation(
|
||||
lambda: wallet.melt(
|
||||
proofs=proofs,
|
||||
invoice=bolt11_invoice,
|
||||
fee_reserve_sat=melt_quote_resp.fee_reserve,
|
||||
quote_id=melt_quote_resp.quote,
|
||||
),
|
||||
op_name="lnurl_melt",
|
||||
mint_url=str(wallet.url),
|
||||
retry_timeouts=False,
|
||||
)
|
||||
except Exception as error:
|
||||
if is_mint_rate_limited(error):
|
||||
# Cooldown failures happen before dispatch, and HTTP 429 means the
|
||||
# mint rejected the request. Neither outcome may keep proofs
|
||||
# reserved as though a Lightning payment could still settle.
|
||||
await wallet.set_reserved_for_send(proofs, reserved=False)
|
||||
raise
|
||||
if not isinstance(error, MINT_TRANSPORT_EXCEPTIONS):
|
||||
raise
|
||||
melt_response = None
|
||||
melt_error: BaseException | None = error
|
||||
else:
|
||||
melt_error = None
|
||||
|
||||
if getattr(melt_response, "state", None) == MeltQuoteState.paid:
|
||||
return final_amount
|
||||
|
||||
try:
|
||||
quote = await run_mint_operation(
|
||||
lambda: wallet.get_melt_quote(melt_quote_resp.quote),
|
||||
op_name="reconcile_lnurl_melt_quote",
|
||||
mint_url=str(wallet.url),
|
||||
retry_timeouts=False,
|
||||
)
|
||||
except Exception as reconciliation_error:
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
"Melt outcome is ambiguous; quote reconciliation failed and proofs "
|
||||
"must not be retried"
|
||||
) from reconciliation_error
|
||||
|
||||
if quote is not None and quote.state == MeltQuoteState.paid:
|
||||
return final_amount
|
||||
|
||||
state = getattr(getattr(quote, "state", None), "value", "unknown")
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
"Melt outcome is ambiguous; proofs must not be retried "
|
||||
f"(quote_state={state})"
|
||||
) from melt_error
|
||||
|
||||
+62
-32
@@ -85,6 +85,30 @@ class Model(BaseModel):
|
||||
return hash(self.id)
|
||||
|
||||
|
||||
def litellm_cost_entry(model_id: str) -> dict | None:
|
||||
"""Look up ``model_id`` in litellm's bundled cost map.
|
||||
|
||||
litellm ships per-model USD rates keyed by the exact OpenRouter id
|
||||
(``deepseek/deepseek-chat``) or the bare model name (``gpt-4o``,
|
||||
``claude-sonnet-4-5``), so both spellings are tried. Keys are lowercase, so
|
||||
a mixed-case upstream id (``deepseek-ai/DeepSeek-V4-Flash``) is retried via
|
||||
a case-insensitive scan. Returns the matched cost dict, or ``None``.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
candidates = (model_id, model_id.split("/", 1)[-1])
|
||||
for key in candidates:
|
||||
info = litellm.model_cost.get(key)
|
||||
if isinstance(info, dict):
|
||||
return info
|
||||
|
||||
lowered = {c.lower() for c in candidates}
|
||||
for key, info in litellm.model_cost.items():
|
||||
if isinstance(key, str) and key.lower() in lowered and isinstance(info, dict):
|
||||
return info
|
||||
return None
|
||||
|
||||
|
||||
def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
|
||||
"""Fill missing cache rates from litellm's bundled cost map.
|
||||
|
||||
@@ -92,12 +116,8 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
|
||||
for many models (most DeepSeek entries, openai/gpt-4o, ...). Without a
|
||||
cache rate, billing falls back to the full input rate, which overcharges
|
||||
cache reads (DeepSeek hits are 10x cheaper) and undercharges Anthropic
|
||||
cache writes (1.25x). litellm ships per-model USD rates keyed by the exact
|
||||
OpenRouter id (deepseek/deepseek-chat) or by the bare model name
|
||||
(gpt-4o, claude-sonnet-4-5), so both spellings are tried. litellm keys are
|
||||
lowercase, but a generic upstream may report a mixed-case id
|
||||
(``deepseek-ai/DeepSeek-V4-Flash``); an exact match is attempted first, then
|
||||
a case-insensitive fallback so such ids still resolve.
|
||||
cache writes (1.25x). The lookup (see ``litellm_cost_entry``) tries both
|
||||
id spellings and a case-insensitive fallback.
|
||||
|
||||
Rates already present (e.g. provided by OpenRouter) are authoritative and
|
||||
never overwritten. Unknown models are returned unchanged.
|
||||
@@ -107,28 +127,7 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
|
||||
if not (needs_read or needs_write):
|
||||
return pricing
|
||||
|
||||
import litellm
|
||||
|
||||
candidates = (model_id, model_id.split("/", 1)[-1])
|
||||
info: dict | None = None
|
||||
for key in candidates:
|
||||
candidate = litellm.model_cost.get(key)
|
||||
if isinstance(candidate, dict):
|
||||
info = candidate
|
||||
break
|
||||
if info is None:
|
||||
# Case-insensitive fallback: a mixed-case upstream id (e.g.
|
||||
# ``deepseek-ai/DeepSeek-V4-Flash``) won't match litellm's lowercase
|
||||
# keys exactly. Build a lowercased index once and retry.
|
||||
lowered = {c.lower() for c in candidates}
|
||||
for key, candidate in litellm.model_cost.items():
|
||||
if (
|
||||
isinstance(key, str)
|
||||
and key.lower() in lowered
|
||||
and isinstance(candidate, dict)
|
||||
):
|
||||
info = candidate
|
||||
break
|
||||
info = litellm_cost_entry(model_id)
|
||||
if info is None:
|
||||
return pricing
|
||||
|
||||
@@ -456,7 +455,9 @@ async def _update_sats_pricing_once() -> None:
|
||||
for m in upstream.get_cached_models()
|
||||
]
|
||||
upstream._models_cache = updated_models
|
||||
upstream._models_by_id = {m.forwarded_model_id or m.id: m for m in updated_models}
|
||||
upstream._models_by_id = {
|
||||
m.forwarded_model_id or m.id: m for m in updated_models
|
||||
}
|
||||
updated_count += len(updated_models)
|
||||
|
||||
if updated_count > 0:
|
||||
@@ -511,9 +512,7 @@ class ModelTestRequest(V2BaseModel):
|
||||
request_data: dict
|
||||
|
||||
|
||||
@models_router.post(
|
||||
"/api/models/test", dependencies=[Depends(_require_admin_api)]
|
||||
)
|
||||
@models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)])
|
||||
async def test_model(
|
||||
payload: ModelTestRequest,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
@@ -596,6 +595,37 @@ async def test_model(
|
||||
}
|
||||
|
||||
|
||||
@models_router.get("/v1/models/paths")
|
||||
@models_router.get("/v1/models/paths/", include_in_schema=False)
|
||||
async def model_paths() -> dict:
|
||||
"""All models with every upstream provider path they are reachable through."""
|
||||
from ..upstream.model_paths import get_all_model_paths
|
||||
|
||||
return await get_all_model_paths()
|
||||
|
||||
|
||||
@models_router.get("/v1/models/paths/model")
|
||||
@models_router.get("/v1/models/paths/model/", include_in_schema=False)
|
||||
async def model_paths_for_model(model_id: str) -> dict:
|
||||
"""Paths for a single model.
|
||||
|
||||
Uses a query parameter (``?model_id=...``) under a fully static route so
|
||||
model ids containing ``/`` (e.g. ``anthropic/claude-opus-4.6``) need no URL
|
||||
encoding and there is no dynamic-route ambiguity.
|
||||
"""
|
||||
from ..proxy import get_unique_models
|
||||
from ..upstream.model_paths import get_paths_for_model
|
||||
|
||||
result = await get_paths_for_model(model_id)
|
||||
if not result["data"]:
|
||||
advertised_ids = {
|
||||
model.forwarded_model_id or model.id for model in get_unique_models()
|
||||
}
|
||||
if model_id not in advertised_ids:
|
||||
raise HTTPException(status_code=404, detail="Model not found")
|
||||
return result
|
||||
|
||||
|
||||
@models_router.get("/v1/models")
|
||||
@models_router.get("/v1/models/", include_in_schema=False)
|
||||
@models_router.get("/models")
|
||||
|
||||
+183
-75
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
@@ -7,7 +8,13 @@ from fastapi.responses import Response, StreamingResponse
|
||||
from sqlmodel import select
|
||||
|
||||
from .algorithm import create_model_mappings
|
||||
from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key
|
||||
from .auth import (
|
||||
ReservationSnapshot,
|
||||
get_reservation_snapshot,
|
||||
pay_for_request,
|
||||
revert_pay_for_request,
|
||||
validate_bearer_key,
|
||||
)
|
||||
from .core import get_logger
|
||||
from .core.db import (
|
||||
ApiKey,
|
||||
@@ -19,7 +26,6 @@ from .core.db import (
|
||||
)
|
||||
from .core.exceptions import UpstreamError
|
||||
from .core.not_found import build_not_found_response
|
||||
from .core.settings import settings
|
||||
from .payment.helpers import (
|
||||
calculate_discounted_max_cost,
|
||||
check_token_balance,
|
||||
@@ -37,13 +43,19 @@ 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)
|
||||
|
||||
|
||||
async def _finish_read_transaction(session: AsyncSession) -> None:
|
||||
"""Release a read transaction without assuming a particular session mock."""
|
||||
commit_result = session.commit()
|
||||
if inspect.isawaitable(commit_result):
|
||||
await commit_result
|
||||
|
||||
|
||||
async def initialize_upstreams() -> None:
|
||||
"""Initialize upstream providers from database during application startup."""
|
||||
global _upstreams
|
||||
@@ -72,32 +84,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]:
|
||||
@@ -106,8 +130,13 @@ def get_unique_models() -> list[Model]:
|
||||
|
||||
|
||||
def _is_tinfoil_attestation_path(path: str) -> bool:
|
||||
"""Return True for Tinfoil attestation-bundle proxy paths."""
|
||||
return path in {"attestation", "tee/attestation"}
|
||||
"""Return True for exact Tinfoil attestation routes, with optional slash."""
|
||||
return path in {
|
||||
"attestation",
|
||||
"attestation/",
|
||||
"tee/attestation",
|
||||
"tee/attestation/",
|
||||
}
|
||||
|
||||
|
||||
def _select_unauthenticated_get_upstreams(
|
||||
@@ -132,7 +161,7 @@ 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
|
||||
@@ -142,24 +171,38 @@ 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,
|
||||
)
|
||||
|
||||
# Keep model-path discovery in sync with admin mutations: disabling or
|
||||
# deleting a provider must stop advertising its paths immediately rather
|
||||
# than after the next timed refresh.
|
||||
from .upstream.model_paths import prune_model_paths_for_inactive_providers
|
||||
|
||||
try:
|
||||
await prune_model_paths_for_inactive_providers()
|
||||
except Exception as e: # noqa: BLE001 - discovery sync must not break routing
|
||||
logger.warning(
|
||||
"Failed to prune model paths for inactive providers",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
|
||||
async def refresh_model_maps_periodically() -> None:
|
||||
"""Background task to refresh model maps every minute."""
|
||||
@@ -197,6 +240,20 @@ _API_PATH_PREFIXES = (
|
||||
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
|
||||
async def proxy(
|
||||
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
||||
) -> Response | StreamingResponse:
|
||||
"""Run proxy setup in a short request session, never across response streaming."""
|
||||
try:
|
||||
return await _proxy(request, path, session)
|
||||
finally:
|
||||
# FastAPI yield dependencies normally close after the response body is
|
||||
# sent. Close explicitly so a long stream cannot retain DB resources.
|
||||
close_result = session.close()
|
||||
if inspect.isawaitable(close_result):
|
||||
await close_result
|
||||
|
||||
|
||||
async def _proxy(
|
||||
request: Request, path: str, session: AsyncSession
|
||||
) -> Response | StreamingResponse:
|
||||
# GET requests must hit a known API prefix; otherwise return a 404 (HTML
|
||||
# for browsers, JSON for API clients). POST requests are always forwarded
|
||||
@@ -233,13 +290,10 @@ async def proxy(
|
||||
else:
|
||||
model_id = request_body_dict.get("model", "unknown")
|
||||
|
||||
# /tee/* and /attestation GET requests don't map to models — forward
|
||||
# without model/cost/auth lookups. Tinfoil attestation paths are routed
|
||||
# only to Tinfoil providers so an unrelated upstream's 404 cannot
|
||||
# short-circuit before the attestation proxy is tried.
|
||||
if request.method == "GET" and (
|
||||
path.startswith("tee/") or path.startswith("attestation")
|
||||
):
|
||||
# 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(
|
||||
@@ -254,7 +308,10 @@ async def proxy(
|
||||
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(selected_upstreams) - 1:
|
||||
if (
|
||||
response.status_code in [502, 429]
|
||||
and i < len(selected_upstreams) - 1
|
||||
):
|
||||
logger.warning(
|
||||
"Upstream %s returned %s for unauthenticated GET %s, trying next",
|
||||
upstream.provider_type,
|
||||
@@ -277,25 +334,20 @@ async def proxy(
|
||||
"upstream_error", "All upstreams failed", 502, request=request
|
||||
)
|
||||
|
||||
model_obj = get_model_instance(model_id)
|
||||
candidates = get_candidates(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:
|
||||
upstreams = [upstream for upstream in upstreams if upstream.supports_ehbp]
|
||||
if not upstreams:
|
||||
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}'",
|
||||
@@ -303,9 +355,10 @@ async def proxy(
|
||||
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
|
||||
@@ -313,14 +366,12 @@ async def proxy(
|
||||
max_cost_for_model = await calculate_discounted_max_cost(
|
||||
_max_cost_for_model, request_body_dict, model_obj=model_obj
|
||||
)
|
||||
# Ensure max_cost_for_model is at least the minimum allowed request cost
|
||||
max_cost_for_model = max(max_cost_for_model, settings.min_request_msat)
|
||||
|
||||
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
||||
|
||||
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_ehbp:
|
||||
if not upstream.supports_ehbp:
|
||||
@@ -358,7 +409,7 @@ async def proxy(
|
||||
"status_code": e.status_code,
|
||||
},
|
||||
)
|
||||
if i == len(upstreams) - 1:
|
||||
if i == len(candidates) - 1:
|
||||
last_error = e
|
||||
continue
|
||||
|
||||
@@ -385,12 +436,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"):
|
||||
@@ -420,22 +471,56 @@ 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
|
||||
)
|
||||
|
||||
reservation_snapshot: ReservationSnapshot | None = None
|
||||
if is_ehbp or request_body_dict:
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
# Snapshot validation performs SELECTs after pay_for_request commits.
|
||||
# End that read transaction before waiting on upstream response headers.
|
||||
await _finish_read_transaction(session)
|
||||
|
||||
# Tracks request params already removed in response to upstream rejections,
|
||||
# shared across providers so a stripped param stays stripped on failover and
|
||||
# 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
|
||||
)
|
||||
if candidate_max > max_cost_for_model:
|
||||
await revert_pay_for_request(
|
||||
key, session, max_cost_for_model, reservation_snapshot
|
||||
)
|
||||
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)
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
await _finish_read_transaction(session)
|
||||
continue
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
await _finish_read_transaction(session)
|
||||
max_cost_for_model = candidate_max
|
||||
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
|
||||
try:
|
||||
@@ -462,6 +547,7 @@ async def proxy(
|
||||
max_cost_for_model=max_cost_for_model,
|
||||
session=session,
|
||||
model_obj=model_obj,
|
||||
reservation_snapshot=reservation_snapshot,
|
||||
)
|
||||
elif is_responses_api:
|
||||
response = await upstream.forward_responses_request(
|
||||
@@ -473,6 +559,7 @@ async def proxy(
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
reservation_snapshot,
|
||||
)
|
||||
else:
|
||||
response = await upstream.forward_request(
|
||||
@@ -484,6 +571,7 @@ async def proxy(
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
reservation_snapshot,
|
||||
)
|
||||
except UpstreamError:
|
||||
# Let the outer UpstreamError handler manage retry/revert
|
||||
@@ -500,7 +588,9 @@ async def proxy(
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
},
|
||||
)
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
await revert_pay_for_request(
|
||||
key, session, max_cost_for_model, reservation_snapshot
|
||||
)
|
||||
raise
|
||||
|
||||
# Reactive recovery: some models reject one specific request
|
||||
@@ -536,7 +626,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"):
|
||||
@@ -569,7 +659,9 @@ async def proxy(
|
||||
continue
|
||||
|
||||
# 4xx error (user error), or other non-retryable error, or last provider failed
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
await revert_pay_for_request(
|
||||
key, session, max_cost_for_model, reservation_snapshot
|
||||
)
|
||||
logger.warning(
|
||||
"Upstream request failed, revert payment "
|
||||
"(provider=%s model=%s status=%s path=%s)",
|
||||
@@ -601,8 +693,10 @@ async def proxy(
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
},
|
||||
)
|
||||
await asyncio.shield(
|
||||
revert_pay_for_request(key, session, max_cost_for_model)
|
||||
# The cancellation has been caught, so complete exact cleanup in
|
||||
# this task before the request-scoped session can be torn down.
|
||||
await revert_pay_for_request(
|
||||
key, session, max_cost_for_model, reservation_snapshot
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -616,13 +710,15 @@ 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:
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
if i == len(candidates) - 1:
|
||||
await revert_pay_for_request(
|
||||
key, session, max_cost_for_model, reservation_snapshot
|
||||
)
|
||||
return create_upstream_error_response(e, request)
|
||||
|
||||
# Otherwise loop continues to next provider
|
||||
@@ -707,17 +803,29 @@ async def get_bearer_token_key(
|
||||
},
|
||||
)
|
||||
return key
|
||||
except Exception as e:
|
||||
key_preview = bearer_key[:20] + "..." if len(bearer_key) > 20 else bearer_key
|
||||
logger.error(
|
||||
f"Bearer token validation failed: {type(e).__name__}: {e} path={path} model={model_id!r} min_cost={min_cost} key={key_preview!r}",
|
||||
except HTTPException as error:
|
||||
detail: dict[str, Any] = error.detail if isinstance(error.detail, dict) else {}
|
||||
raw_error = detail.get("error")
|
||||
error_info = raw_error if isinstance(raw_error, dict) else {}
|
||||
logger.warning(
|
||||
"Bearer token rejected",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"status_code": error.status_code,
|
||||
"error_code": error_info.get("code"),
|
||||
"path": path,
|
||||
"model_id": model_id,
|
||||
"min_cost_msat": min_cost,
|
||||
"bearer_key_preview": key_preview,
|
||||
"required_msat": min_cost,
|
||||
},
|
||||
)
|
||||
raise
|
||||
except Exception as error:
|
||||
logger.exception(
|
||||
"Bearer token validation failed",
|
||||
extra={
|
||||
"error_type": type(error).__name__,
|
||||
"path": path,
|
||||
"model_id": model_id,
|
||||
"required_msat": min_cost,
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
+1009
-12
File diff suppressed because it is too large
Load Diff
+550
-195
File diff suppressed because it is too large
Load Diff
+161
-20
@@ -15,7 +15,11 @@ from sqlmodel import col, update
|
||||
|
||||
from ..auth import (
|
||||
ROUTSTR_FEE_PERCENT,
|
||||
ReservationSnapshot,
|
||||
_claim_reservation_for_charge,
|
||||
_validate_reservation_snapshot,
|
||||
get_billing_key,
|
||||
get_reservation_snapshot,
|
||||
payments_logger,
|
||||
)
|
||||
from ..core import get_logger
|
||||
@@ -23,7 +27,9 @@ from ..core.db import (
|
||||
ApiKey,
|
||||
AsyncSession,
|
||||
accumulate_routstr_fee,
|
||||
store_cashu_transaction,
|
||||
)
|
||||
from ..core.db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from ..core.exceptions import UpstreamError
|
||||
from ..core.settings import settings
|
||||
@@ -50,6 +56,14 @@ _TINFOIL_PROVIDER_TYPE = "tinfoil"
|
||||
_TINFOIL_ALLOWED_ENCLAVE_HOST_SUFFIX = ".tinfoil.sh"
|
||||
_TINFOIL_ALLOWED_ENCLAVE_HOSTS = frozenset({"tinfoil.sh"})
|
||||
|
||||
|
||||
def _normalize_upstream_model_id(model_id: str | None) -> str:
|
||||
"""Normalize casing and whitespace for upstream identity comparisons."""
|
||||
if not model_id:
|
||||
return ""
|
||||
return model_id.strip().lower()
|
||||
|
||||
|
||||
# Headers that must not be forwarded to the upstream enclave.
|
||||
_PROXY_ONLY_HEADERS = frozenset(
|
||||
{
|
||||
@@ -63,30 +77,48 @@ _PROXY_ONLY_HEADERS = frozenset(
|
||||
def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None:
|
||||
"""Parse ``X-Tinfoil-Usage-Metrics`` into an OpenAI-style usage dict.
|
||||
|
||||
The header format is ``prompt=<n>,completion=<n>,total=<n>``. Returns a dict
|
||||
like ``{"prompt_tokens": n, "completion_tokens": n}`` suitable for
|
||||
:func:`calculate_cost`, or ``None`` when the header is absent or malformed.
|
||||
The header format is::
|
||||
|
||||
prompt=<n>,completion=<n>,total=<n>[,model=<name>]
|
||||
|
||||
The ``model`` field (added in tinfoilsh/confidential-model-router PR #385)
|
||||
is extracted as a string and included in the returned dict under the
|
||||
``"model"`` key so callers can compare the served model against the
|
||||
requested one and adjust pricing.
|
||||
|
||||
Returns a dict like ``{"prompt_tokens": n, "completion_tokens": n,
|
||||
"model": "<name>"}`` suitable for :func:`calculate_cost` (which ignores
|
||||
the extra ``model`` key in the usage sub-dict), or ``None`` when the
|
||||
header is absent or malformed.
|
||||
"""
|
||||
if not header_value:
|
||||
return None
|
||||
parts: dict[str, int] = {}
|
||||
model: str | None = None
|
||||
for item in header_value.split(","):
|
||||
key, sep, value = item.partition("=")
|
||||
if not sep:
|
||||
continue
|
||||
key = key.strip()
|
||||
value = value.strip()
|
||||
if key == "model":
|
||||
model = value
|
||||
continue
|
||||
try:
|
||||
parts[key.strip()] = int(value.strip())
|
||||
parts[key] = int(value)
|
||||
except (ValueError, TypeError):
|
||||
continue
|
||||
prompt = parts.get("prompt")
|
||||
completion = parts.get("completion")
|
||||
if prompt is not None and completion is not None:
|
||||
result: dict[str, int] = {
|
||||
result: dict[str, int | str] = {
|
||||
"prompt_tokens": prompt,
|
||||
"completion_tokens": completion,
|
||||
}
|
||||
if "total" in parts:
|
||||
result["total_tokens"] = parts["total"]
|
||||
if model:
|
||||
result["model"] = model
|
||||
return result
|
||||
logger.warning(
|
||||
"Failed to parse X-Tinfoil-Usage-Metrics header",
|
||||
@@ -242,9 +274,15 @@ def _build_cost_info(
|
||||
output_tokens: int = 0,
|
||||
input_msats: int = 0,
|
||||
output_msats: int = 0,
|
||||
actual_model: str | None = None,
|
||||
) -> dict:
|
||||
"""Build a cost-info dict with token counts and per-token-type costs."""
|
||||
return {
|
||||
"""Build a cost-info dict with token counts and per-token-type costs.
|
||||
|
||||
When ``actual_model`` is set (the served model differs from the requested
|
||||
one), it is included in the returned dict so callers can use it for billing
|
||||
finalization and logging.
|
||||
"""
|
||||
result: dict[str, int | str | None] = {
|
||||
"total_msats": total_msats,
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
@@ -252,6 +290,9 @@ def _build_cost_info(
|
||||
"input_msats": input_msats,
|
||||
"output_msats": output_msats,
|
||||
}
|
||||
if actual_model:
|
||||
result["actual_model"] = actual_model
|
||||
return result
|
||||
|
||||
|
||||
def _inject_cost_response_headers(
|
||||
@@ -280,48 +321,118 @@ async def _compute_ehbp_actual_cost(
|
||||
max_cost_for_model]`` so the refund never exceeds the reservation and is
|
||||
never zero.
|
||||
|
||||
When the usage-metrics header includes ``model=<name>`` and it differs
|
||||
from ``model_obj.id``, the actual served model's pricing is used for the
|
||||
cost calculation. The returned dict includes an ``"actual_model"`` key
|
||||
in that case so callers can use it for billing finalization.
|
||||
|
||||
Returns a dict with ``total_msats``, ``input_tokens``, ``output_tokens``,
|
||||
``total_tokens``, ``input_msats``, and ``output_msats``.
|
||||
``total_tokens``, ``input_msats``, and ``output_msats`` (and optionally
|
||||
``actual_model``).
|
||||
"""
|
||||
usage_dict = parse_tinfoil_usage_metrics(usage_header)
|
||||
if usage_dict is None:
|
||||
return _build_cost_info(max_cost_for_model)
|
||||
|
||||
# The enclave may serve a different model than the one requested (e.g.
|
||||
# due to failover). The usage-metrics header's ``model=<name>`` carries
|
||||
# the actual upstream model ID (e.g. ``glm-5-2``), which may differ from
|
||||
# the client-facing ``model_obj.id`` (e.g. ``tinfoil-glm-5-2``) even when
|
||||
# the correct model was served — the alias is resolved through
|
||||
# ``model_obj.forwarded_model_id``. Only when the served model differs
|
||||
# from the expected upstream ID do we treat it as a real mismatch and
|
||||
# look up the actual model's pricing.
|
||||
actual_model: str | None = usage_dict.pop("model", None) # type: ignore[arg-type]
|
||||
pricing_model_id = model_obj.id
|
||||
expected_upstream_model = model_obj.forwarded_model_id or model_obj.id
|
||||
expected_identity = _normalize_upstream_model_id(expected_upstream_model)
|
||||
served_identity = _normalize_upstream_model_id(actual_model)
|
||||
|
||||
# Ignore casing and surrounding whitespace when comparing the model
|
||||
# reported by the enclave with the expected upstream model. Version
|
||||
# suffixes remain part of the identity because a configured
|
||||
# ``forwarded_model_id`` may intentionally include one.
|
||||
if actual_model and served_identity != expected_identity:
|
||||
from ..proxy import get_model_instance
|
||||
|
||||
# ``forwarded_model_id`` values are registered as routable aliases in
|
||||
# the global model map. The resolved object can belong to a different
|
||||
# provider and therefore have a different client-facing ``id`` while
|
||||
# still representing the same upstream model.
|
||||
actual_model_obj = get_model_instance(actual_model)
|
||||
if actual_model_obj is None:
|
||||
logger.warning(
|
||||
"EHBP served model not found in registry, falling back "
|
||||
"to requested model for pricing",
|
||||
extra={
|
||||
"requested_model": model_obj.id,
|
||||
"expected_upstream_model": expected_upstream_model,
|
||||
"actual_model": actual_model,
|
||||
},
|
||||
)
|
||||
actual_model = None
|
||||
else:
|
||||
resolved_upstream_model = (
|
||||
actual_model_obj.forwarded_model_id or actual_model_obj.id
|
||||
)
|
||||
resolved_identity = _normalize_upstream_model_id(
|
||||
resolved_upstream_model
|
||||
)
|
||||
if resolved_identity != expected_identity:
|
||||
logger.info(
|
||||
"EHBP served model differs from requested, using actual "
|
||||
"model for pricing",
|
||||
extra={
|
||||
"requested_model": model_obj.id,
|
||||
"expected_upstream_model": expected_upstream_model,
|
||||
"actual_model": actual_model,
|
||||
"resolved_upstream_model": resolved_upstream_model,
|
||||
},
|
||||
)
|
||||
pricing_model_id = actual_model_obj.id
|
||||
else:
|
||||
# A different registry/client alias resolved to the same
|
||||
# upstream model; retain the requested model's pricing.
|
||||
actual_model = None
|
||||
else:
|
||||
# Models match or no model in header — use requested model's pricing.
|
||||
actual_model = None
|
||||
|
||||
try:
|
||||
cost = await calculate_cost(
|
||||
{"model": model_obj.id, "usage": usage_dict},
|
||||
{"model": pricing_model_id, "usage": usage_dict},
|
||||
max_cost_for_model,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"EHBP usage cost calculation failed, falling back to max cost",
|
||||
extra={
|
||||
"model": model_obj.id,
|
||||
"model": pricing_model_id,
|
||||
"error": str(e),
|
||||
"usage": usage_dict,
|
||||
},
|
||||
)
|
||||
return _build_cost_info(max_cost_for_model)
|
||||
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
|
||||
|
||||
if isinstance(cost, MaxCostData):
|
||||
logger.warning(
|
||||
"EHBP calculate_cost returned MaxCostData (no model pricing), "
|
||||
"falling back to max cost",
|
||||
extra={
|
||||
"model": model_obj.id,
|
||||
"model": pricing_model_id,
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
"usage": usage_dict,
|
||||
"cost_total_msats": cost.total_msats,
|
||||
},
|
||||
)
|
||||
return _build_cost_info(max_cost_for_model)
|
||||
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
|
||||
if isinstance(cost, CostData):
|
||||
actual = max(int(cost.total_msats), int(settings.min_request_msat))
|
||||
clamped = min(actual, max_cost_for_model)
|
||||
logger.info(
|
||||
"EHBP actual cost computed from usage metrics",
|
||||
extra={
|
||||
"model": model_obj.id,
|
||||
"model": pricing_model_id,
|
||||
"usage": usage_dict,
|
||||
"cost_total_msats": cost.total_msats,
|
||||
"clamped_msats": clamped,
|
||||
@@ -334,16 +445,17 @@ async def _compute_ehbp_actual_cost(
|
||||
output_tokens=cost.output_tokens,
|
||||
input_msats=cost.input_msats,
|
||||
output_msats=cost.output_msats,
|
||||
actual_model=actual_model,
|
||||
)
|
||||
# CostDataError
|
||||
logger.warning(
|
||||
"EHBP usage cost calculation error, falling back to max cost",
|
||||
extra={
|
||||
"model": model_obj.id,
|
||||
"model": pricing_model_id,
|
||||
"error": getattr(cost, "message", str(cost)),
|
||||
},
|
||||
)
|
||||
return _build_cost_info(max_cost_for_model)
|
||||
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
|
||||
|
||||
|
||||
def _extract_usage_from_response(
|
||||
@@ -394,8 +506,14 @@ async def finalize_ehbp_actual_cost_payment(
|
||||
reserved_cost_for_model: int,
|
||||
model_id: str,
|
||||
cost_info: dict,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
) -> None:
|
||||
"""Finalize an EHBP bearer request using clamped provider usage metrics."""
|
||||
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
||||
await _validate_reservation_snapshot(key, reservation, session)
|
||||
if not await _claim_reservation_for_charge(reservation, session):
|
||||
return
|
||||
reserved_cost_for_model = reservation.reserved_msats
|
||||
billing_key = await get_billing_key(key, session)
|
||||
key_hash = key.hashed_key
|
||||
billing_key_hash = billing_key.hashed_key
|
||||
@@ -498,6 +616,7 @@ async def finalize_ehbp_max_cost_payment(
|
||||
session: AsyncSession,
|
||||
max_cost_for_model: int,
|
||||
model_id: str,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
) -> None:
|
||||
"""Finalize an EHBP bearer request by charging the reserved max cost.
|
||||
|
||||
@@ -505,6 +624,11 @@ async def finalize_ehbp_max_cost_payment(
|
||||
normal completion handlers, this intentionally charges the pre-reserved max
|
||||
cost and releases the reservation.
|
||||
"""
|
||||
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
||||
await _validate_reservation_snapshot(key, reservation, session)
|
||||
if not await _claim_reservation_for_charge(reservation, session):
|
||||
return
|
||||
max_cost_for_model = reservation.reserved_msats
|
||||
billing_key = await get_billing_key(key, session)
|
||||
key_hash = key.hashed_key
|
||||
billing_key_hash = billing_key.hashed_key
|
||||
@@ -658,6 +782,7 @@ async def forward_ehbp_request(
|
||||
max_cost_for_model: int,
|
||||
session: AsyncSession,
|
||||
model_obj: Model,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
) -> Response | StreamingResponse:
|
||||
"""Forward an EHBP bearer-auth request and finalize billing.
|
||||
|
||||
@@ -771,8 +896,16 @@ async def forward_ehbp_request(
|
||||
cost_info = await _compute_ehbp_actual_cost(
|
||||
usage_header, model_obj, max_cost_for_model
|
||||
)
|
||||
# Use the actual served model for billing when it differs from
|
||||
# the requested model.
|
||||
billing_model = cost_info.pop("actual_model", None) or model_obj.id
|
||||
await finalize_ehbp_actual_cost_payment(
|
||||
key, session, max_cost_for_model, model_obj.id, cost_info
|
||||
key,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
billing_model,
|
||||
cost_info,
|
||||
reservation_snapshot,
|
||||
)
|
||||
cost_data = {**cost_info, "total_usd": 0.0}
|
||||
else:
|
||||
@@ -786,7 +919,11 @@ async def forward_ehbp_request(
|
||||
},
|
||||
)
|
||||
await finalize_ehbp_max_cost_payment(
|
||||
key, session, max_cost_for_model, model_obj.id
|
||||
key,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj.id,
|
||||
reservation_snapshot,
|
||||
)
|
||||
cost_data = {
|
||||
"total_msats": max_cost_for_model,
|
||||
@@ -972,11 +1109,15 @@ async def forward_ehbp_x_cashu_request(
|
||||
usage_header, model_obj, max_cost_for_model
|
||||
)
|
||||
actual_cost_msats = cost_info["total_msats"]
|
||||
actual_model = cost_info.get("actual_model")
|
||||
billing_model = actual_model or model_obj.id
|
||||
refund_amount = amount - _msats_to_unit_amount(actual_cost_msats, unit)
|
||||
logger.info(
|
||||
"EHBP X-Cashu refund computed",
|
||||
extra={
|
||||
"model": model_obj.id,
|
||||
"model": billing_model,
|
||||
"requested_model": model_obj.id,
|
||||
"actual_model": actual_model,
|
||||
"redeemed_amount": amount,
|
||||
"actual_cost_msats": actual_cost_msats,
|
||||
"refund_amount": refund_amount,
|
||||
|
||||
+91
-44
@@ -5,6 +5,12 @@ from typing import TYPE_CHECKING
|
||||
import httpx
|
||||
|
||||
from .base import BaseUpstreamProvider
|
||||
from .pricing_resolver import (
|
||||
FallbackPricingResolver,
|
||||
ResolvedPricing,
|
||||
_as_float,
|
||||
estimate_context_length,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
@@ -64,6 +70,40 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def _native_pricing(
|
||||
self, model_id: str, model_spec: dict
|
||||
) -> ResolvedPricing | None:
|
||||
"""Read pricing/metadata from Venice's bespoke ``model_spec`` schema.
|
||||
|
||||
Returns ``None`` when the upstream reported no *usable* native price —
|
||||
absent, non-numeric, negative, or both-zero — so the caller falls
|
||||
through to the shared resolution chain instead of fabricating a number
|
||||
or trusting a bogus one. This mirrors the money-safety guards the
|
||||
litellm and OpenRouter rungs already apply: a both-zero price would
|
||||
serve the model free, a negative one would credit the caller, and a
|
||||
non-numeric string would otherwise throw and drop the whole catalog.
|
||||
"""
|
||||
pricing_info = model_spec.get("pricing", {})
|
||||
input_usd = _as_float(pricing_info.get("input", {}).get("usd"))
|
||||
output_usd = _as_float(pricing_info.get("output", {}).get("usd"))
|
||||
if input_usd is None or output_usd is None:
|
||||
return None
|
||||
if input_usd < 0 or output_usd < 0 or (input_usd == 0 and output_usd == 0):
|
||||
return None
|
||||
|
||||
capabilities = model_spec.get("capabilities", {})
|
||||
input_modalities = ["text"]
|
||||
if capabilities.get("supportsVision", False):
|
||||
input_modalities.append("image")
|
||||
|
||||
return ResolvedPricing(
|
||||
prompt=input_usd / 1_000_000,
|
||||
completion=output_usd / 1_000_000,
|
||||
context_length=model_spec.get("availableContextTokens"),
|
||||
source="native",
|
||||
input_modalities=input_modalities,
|
||||
)
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch models from upstream API using /models endpoint."""
|
||||
from ..payment.models import Architecture, Model, Pricing, TopProvider
|
||||
@@ -78,6 +118,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
resolver = FallbackPricingResolver()
|
||||
models_list = []
|
||||
for model_data in data.get("data", []):
|
||||
model_id = model_data.get("id", "")
|
||||
@@ -89,41 +130,44 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
|
||||
owned_by = model_data.get("owned_by", "unknown")
|
||||
model_spec = model_data.get("model_spec", {})
|
||||
|
||||
context_length = 4096
|
||||
if model_spec.get("availableContextTokens"):
|
||||
context_length = model_spec["availableContextTokens"]
|
||||
elif any(
|
||||
pattern in model_id.lower() for pattern in ["32k", "32000"]
|
||||
):
|
||||
context_length = 32768
|
||||
elif any(
|
||||
pattern in model_id.lower() for pattern in ["16k", "16000"]
|
||||
):
|
||||
context_length = 16384
|
||||
elif any(pattern in model_id.lower() for pattern in ["8k", "8000"]):
|
||||
context_length = 8192
|
||||
elif "gpt-4" in model_id.lower():
|
||||
context_length = 8192
|
||||
elif "claude" in model_id.lower():
|
||||
context_length = 200000
|
||||
resolved = self._native_pricing(model_id, model_spec)
|
||||
if resolved is None:
|
||||
resolved = await resolver.resolve(model_id)
|
||||
|
||||
pricing_info = model_spec.get("pricing", {})
|
||||
input_pricing = pricing_info.get("input", {})
|
||||
output_pricing = pricing_info.get("output", {})
|
||||
if resolved is None:
|
||||
# Fail closed: never invent a price. Import the model
|
||||
# disabled with a warning so the operator can price it
|
||||
# (the admin UI surfaces disabled remote models).
|
||||
logger.warning(
|
||||
f"No pricing source resolved for '{model_id}' from "
|
||||
f"{self.upstream_name}; importing it disabled",
|
||||
extra={"model_id": model_id, "base_url": self.base_url},
|
||||
)
|
||||
resolved = ResolvedPricing(
|
||||
prompt=0.0,
|
||||
completion=0.0,
|
||||
context_length=None,
|
||||
source="unresolved",
|
||||
)
|
||||
enabled = False
|
||||
else:
|
||||
enabled = True
|
||||
|
||||
prompt_price = input_pricing.get("usd", 0.001) / 1000000
|
||||
completion_price = output_pricing.get("usd", 0.001) / 1000000
|
||||
# Prefer the source's own modality string (OpenRouter ships
|
||||
# one, e.g. "text+image->text"); otherwise derive it from the
|
||||
# captured input/output modalities in the same "in->out" shape
|
||||
# rather than flattening vision models to "text->text".
|
||||
modality = resolved.modality or (
|
||||
f"{'+'.join(resolved.input_modalities)}"
|
||||
f"->{'+'.join(resolved.output_modalities)}"
|
||||
)
|
||||
|
||||
capabilities = model_spec.get("capabilities", {})
|
||||
input_modalities = ["text"]
|
||||
output_modalities = ["text"]
|
||||
|
||||
if capabilities.get("supportsVision", False):
|
||||
input_modalities.append("image")
|
||||
|
||||
modality = "text"
|
||||
if capabilities.get("supportsVision", False):
|
||||
modality = "text->text"
|
||||
# A source can carry a price but no context (e.g. a litellm
|
||||
# entry missing max_input_tokens); fall back to an id-based
|
||||
# estimate so we never persist a zero-length window.
|
||||
context_length = resolved.context_length or estimate_context_length(
|
||||
model_id
|
||||
)
|
||||
|
||||
spec_name = model_spec.get("name", model_name)
|
||||
description = f"{spec_name}"
|
||||
@@ -139,30 +183,33 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
|
||||
context_length=context_length,
|
||||
architecture=Architecture(
|
||||
modality=modality,
|
||||
input_modalities=input_modalities,
|
||||
output_modalities=output_modalities,
|
||||
tokenizer="unknown",
|
||||
instruct_type=None,
|
||||
input_modalities=resolved.input_modalities,
|
||||
output_modalities=resolved.output_modalities,
|
||||
tokenizer=resolved.tokenizer,
|
||||
instruct_type=resolved.instruct_type,
|
||||
),
|
||||
pricing=Pricing(
|
||||
prompt=prompt_price,
|
||||
completion=completion_price,
|
||||
prompt=resolved.prompt,
|
||||
completion=resolved.completion,
|
||||
request=0.0,
|
||||
image=0.0,
|
||||
web_search=0.0,
|
||||
internal_reasoning=0.0,
|
||||
max_prompt_cost=0.001,
|
||||
max_completion_cost=0.001,
|
||||
max_cost=0.001,
|
||||
input_cache_read=resolved.input_cache_read,
|
||||
input_cache_write=resolved.input_cache_write,
|
||||
),
|
||||
sats_pricing=None,
|
||||
per_request_limits=None,
|
||||
top_provider=TopProvider(
|
||||
context_length=context_length,
|
||||
max_completion_tokens=context_length // 2,
|
||||
is_moderated=False,
|
||||
max_completion_tokens=(
|
||||
resolved.max_completion_tokens
|
||||
if resolved.max_completion_tokens is not None
|
||||
else context_length // 2
|
||||
),
|
||||
is_moderated=bool(resolved.is_moderated),
|
||||
),
|
||||
enabled=True,
|
||||
enabled=enabled,
|
||||
upstream_provider_id=None,
|
||||
canonical_slug=None,
|
||||
)
|
||||
|
||||
+19
-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())
|
||||
|
||||
|
||||
@@ -0,0 +1,819 @@
|
||||
"""Model-path discovery service.
|
||||
|
||||
Exposes every selectable upstream route a Routstr model is reachable through.
|
||||
This PR remains discovery-only: request-side routing will consume the opaque
|
||||
selectors in a follow-up.
|
||||
|
||||
A path is a standard percent-encoded query string containing the configured
|
||||
upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter
|
||||
endpoint, its machine-readable tag::
|
||||
|
||||
url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=claude-sonnet-4
|
||||
url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=claude-sonnet-4&endpoint=google-vertex%2Fus
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ipaddress
|
||||
import random
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Callable
|
||||
from urllib.parse import urlencode, urlsplit
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.dialects.sqlite import insert
|
||||
from sqlalchemy.orm import selectinload
|
||||
from sqlmodel import col, delete, select
|
||||
|
||||
from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_session
|
||||
from ..core.logging import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Bound the per-model OpenRouter /endpoints fan-out so a provider with hundreds
|
||||
# of models does not open hundreds of concurrent requests every refresh.
|
||||
_OPENROUTER_CONCURRENCY = 5
|
||||
_OPENROUTER_TIMEOUT_SECONDS = 10.0
|
||||
|
||||
# Rows inserted per statement during persist. Keeps each INSERT bounded while
|
||||
# avoiding per-row round-trips that hold SQLite's write lock for ~1s per cycle.
|
||||
_PERSIST_CHUNK_SIZE = 500
|
||||
|
||||
# Admin mutations enqueue provider IDs here instead of running OpenRouter's
|
||||
# per-model endpoint fan-out inside the request. One worker serializes refreshes
|
||||
# and coalesces repeated mutations for the same provider.
|
||||
_scheduled_provider_refresh_ids: set[int] = set()
|
||||
_scheduled_provider_refresh_task: asyncio.Task[None] | None = None
|
||||
|
||||
# Visibility key used across this module: routing carries the provider
|
||||
# dimension everywhere (ModelRow's primary key is (id, upstream_provider_id)),
|
||||
# so all model-id keyed maps here do too, lowercased like proxy.refresh_model_maps.
|
||||
ModelKey = tuple[str, int]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EndpointIdentity:
|
||||
"""Exact OpenRouter endpoint identity returned by ``/endpoints``."""
|
||||
|
||||
tag: str
|
||||
provider_name: str | None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ConfiguredProviderIdentity:
|
||||
"""Public identity of one configured upstream provider."""
|
||||
|
||||
id: int
|
||||
slug: str
|
||||
provider_type: str
|
||||
base_url: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DiscoveredPath:
|
||||
"""One model route ready for persistence and API serialization."""
|
||||
|
||||
model_id: str
|
||||
path: str
|
||||
provider: ConfiguredProviderIdentity
|
||||
endpoint_tag: str | None = None
|
||||
endpoint_name: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProviderPathSnapshot:
|
||||
"""Refresh result plus model IDs whose prior rows must survive degradation."""
|
||||
|
||||
paths: tuple[DiscoveredPath, ...]
|
||||
preserve_model_ids: frozenset[str] = frozenset()
|
||||
|
||||
|
||||
def public_provider_url(base_url: str) -> str:
|
||||
"""Mask private IP addresses and URLs with explicit ports."""
|
||||
parsed = urlsplit(base_url)
|
||||
try:
|
||||
if parsed.port is not None:
|
||||
return "http://localhost"
|
||||
except ValueError:
|
||||
# An invalid explicit port must not accidentally leak through.
|
||||
return "http://localhost"
|
||||
|
||||
hostname = parsed.hostname
|
||||
if hostname is None:
|
||||
return base_url
|
||||
try:
|
||||
address = ipaddress.ip_address(hostname)
|
||||
except ValueError:
|
||||
return base_url
|
||||
return "http://localhost" if address.is_private else base_url
|
||||
|
||||
|
||||
def encode_model_path(
|
||||
base_url: str,
|
||||
provider_id: int,
|
||||
model_id: str,
|
||||
endpoint_tag: str | None = None,
|
||||
) -> str:
|
||||
"""Encode the complete upstream route selector advertised to clients."""
|
||||
components: list[tuple[str, str | int]] = [
|
||||
("url", base_url),
|
||||
("provider-id", provider_id),
|
||||
("model-id", model_id),
|
||||
]
|
||||
if endpoint_tag:
|
||||
components.append(("endpoint", endpoint_tag))
|
||||
return urlencode(components)
|
||||
|
||||
|
||||
def _make_http_client() -> httpx.AsyncClient:
|
||||
"""Client factory, separated so tests can substitute a mock transport."""
|
||||
return httpx.AsyncClient()
|
||||
|
||||
|
||||
def is_openrouter_base_url(base_url: str | None) -> bool:
|
||||
"""True when ``base_url`` points at OpenRouter.
|
||||
|
||||
Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``:
|
||||
that predicate also returns True for native Anthropic (correct for
|
||||
cache-control, wrong for OpenRouter endpoint discovery). This one keys only
|
||||
on the URL so a ``GenericUpstreamProvider`` aimed at OpenRouter is matched
|
||||
while native Anthropic is not.
|
||||
"""
|
||||
return "openrouter.ai" in (base_url or "")
|
||||
|
||||
|
||||
def exposed_model_id(model: object) -> str:
|
||||
"""Return exactly the ID advertised by ``/v1/models``.
|
||||
|
||||
A forwarded ID is already a public routable alias and must remain intact,
|
||||
including any slash. Without one, ``/v1/models`` exposes the base ID.
|
||||
"""
|
||||
forwarded = getattr(model, "forwarded_model_id", None)
|
||||
if forwarded:
|
||||
return forwarded
|
||||
return public_model_id(getattr(model, "id"))
|
||||
|
||||
|
||||
def public_model_id(model_id: str) -> str:
|
||||
"""Model id exposed by model-path API responses.
|
||||
|
||||
Uses the same rule as ``create_model_mappings.get_base_model_id`` and
|
||||
``resolve_model_alias`` — strip everything before the *first* slash — so
|
||||
the id shown here can be sent back to ``/v1/chat/completions`` verbatim.
|
||||
"""
|
||||
return model_id.split("/", 1)[1] if "/" in model_id else model_id
|
||||
|
||||
|
||||
def openrouter_author_slug(model: object) -> str | None:
|
||||
"""Return a canonical ``author/slug`` for the OpenRouter endpoints API.
|
||||
|
||||
Prefer ``canonical_slug``, then a slash-containing ``id``, then a
|
||||
slash-containing ``forwarded_model_id``. The forwarded id is exactly what
|
||||
the proxy sends upstream for admin-created alias rows (``base.py`` forwards
|
||||
``forwarded_model_id or id``), so it is a valid OpenRouter id when the
|
||||
bare ``id`` is a local alias with no slash.
|
||||
"""
|
||||
canonical = getattr(model, "canonical_slug", None)
|
||||
if canonical and "/" in canonical:
|
||||
return canonical
|
||||
model_id = getattr(model, "id", None)
|
||||
if model_id and "/" in model_id:
|
||||
return model_id
|
||||
forwarded = getattr(model, "forwarded_model_id", None)
|
||||
if forwarded and "/" in forwarded:
|
||||
return forwarded
|
||||
return None
|
||||
|
||||
|
||||
class _RefreshCycleState:
|
||||
"""Per-refresh shared state: fetch dedupe cache and rate-limit latch.
|
||||
|
||||
``endpoint_cache`` dedupes byte-identical ``/endpoints`` fetches when two
|
||||
providers point at the same OpenRouter base URL. ``rate_limited`` latches
|
||||
on the first 429 so the rest of the cycle stops hammering a throttled API;
|
||||
the whole provider result then degrades to "unknown" instead of an empty
|
||||
list, which preserves previously persisted rows.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.endpoint_cache: dict[tuple[str, str], list[EndpointIdentity] | None] = {}
|
||||
self.rate_limited = False
|
||||
|
||||
|
||||
async def _fetch_openrouter_endpoint_subproviders(
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
author_slug: str,
|
||||
semaphore: asyncio.Semaphore,
|
||||
cycle: _RefreshCycleState,
|
||||
) -> list[EndpointIdentity] | None:
|
||||
"""Return exact endpoint identities for one model, or ``None`` when unknown.
|
||||
|
||||
``None`` (not ``[]``) signals a degraded fetch — network failure, rate
|
||||
limit, non-200, or an unparseable payload — so callers can distinguish
|
||||
"this model has no endpoints" from "we could not find out". Failures are
|
||||
logged and swallowed so one model never breaks the whole refresh.
|
||||
"""
|
||||
cache_key = (base_url, author_slug)
|
||||
if cache_key in cycle.endpoint_cache:
|
||||
return cycle.endpoint_cache[cache_key]
|
||||
if cycle.rate_limited:
|
||||
return None
|
||||
|
||||
url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints"
|
||||
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
|
||||
result: list[EndpointIdentity] | None
|
||||
async with semaphore:
|
||||
try:
|
||||
resp = await client.get(
|
||||
url, headers=headers, timeout=_OPENROUTER_TIMEOUT_SECONDS
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 - isolate per-model failures
|
||||
logger.warning(
|
||||
"OpenRouter endpoint discovery request failed",
|
||||
extra={"author_slug": author_slug, "error": str(e)},
|
||||
)
|
||||
cycle.endpoint_cache[cache_key] = None
|
||||
return None
|
||||
|
||||
if resp.status_code == 429:
|
||||
logger.warning(
|
||||
"OpenRouter endpoint discovery rate-limited; aborting cycle",
|
||||
extra={"author_slug": author_slug},
|
||||
)
|
||||
cycle.rate_limited = True
|
||||
cycle.endpoint_cache[cache_key] = None
|
||||
return None
|
||||
if resp.status_code != 200:
|
||||
logger.warning(
|
||||
"OpenRouter endpoint discovery non-200",
|
||||
extra={"author_slug": author_slug, "status_code": resp.status_code},
|
||||
)
|
||||
cycle.endpoint_cache[cache_key] = None
|
||||
return None
|
||||
|
||||
try:
|
||||
payload = resp.json()
|
||||
data = payload.get("data") if isinstance(payload, dict) else None
|
||||
endpoints = data.get("endpoints") if isinstance(data, dict) else None
|
||||
if not isinstance(endpoints, list):
|
||||
raise ValueError("endpoints must be a list")
|
||||
identities: dict[str, EndpointIdentity] = {}
|
||||
for endpoint in endpoints:
|
||||
if not isinstance(endpoint, dict):
|
||||
continue
|
||||
tag = endpoint.get("tag")
|
||||
if not isinstance(tag, str) or not tag.strip():
|
||||
continue
|
||||
provider_name = endpoint.get("provider_name")
|
||||
identities.setdefault(
|
||||
tag,
|
||||
EndpointIdentity(
|
||||
tag=tag,
|
||||
provider_name=provider_name
|
||||
if isinstance(provider_name, str) and provider_name
|
||||
else None,
|
||||
),
|
||||
)
|
||||
if endpoints and not identities:
|
||||
raise ValueError("endpoints contain no usable tags")
|
||||
result = list(identities.values())
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning(
|
||||
"OpenRouter endpoint discovery bad payload",
|
||||
extra={"author_slug": author_slug, "error": str(e)},
|
||||
)
|
||||
result = None
|
||||
|
||||
cycle.endpoint_cache[cache_key] = result
|
||||
return result
|
||||
|
||||
|
||||
async def _load_model_visibility() -> tuple[
|
||||
dict[ModelKey, ModelRow],
|
||||
set[ModelKey],
|
||||
dict[int, ConfiguredProviderIdentity],
|
||||
]:
|
||||
"""Load the same DB model visibility inputs used by routing.
|
||||
|
||||
``refresh_model_maps`` builds routing from enabled providers, enabled DB
|
||||
override rows, and disabled model keys — all keyed on
|
||||
``(model_id.lower(), upstream_provider_id)`` because ``ModelRow``'s primary
|
||||
key is composite and the same id legitimately exists on several providers.
|
||||
Model-path discovery uses the same keying so disabling a model on one
|
||||
provider never hides it on another, and one provider's
|
||||
``forwarded_model_id`` alias is never applied to a different provider.
|
||||
"""
|
||||
async with create_session() as session:
|
||||
query = select(UpstreamProviderRow).options(
|
||||
selectinload(UpstreamProviderRow.models) # type: ignore[arg-type]
|
||||
)
|
||||
provider_rows = (await session.exec(query)).all()
|
||||
|
||||
overrides_by_key: dict[ModelKey, ModelRow] = {}
|
||||
disabled_model_keys: set[ModelKey] = set()
|
||||
provider_identities: dict[int, ConfiguredProviderIdentity] = {}
|
||||
|
||||
for provider in provider_rows:
|
||||
if not provider.enabled or provider.id is None:
|
||||
continue
|
||||
provider_identities[provider.id] = ConfiguredProviderIdentity(
|
||||
id=provider.id,
|
||||
slug=provider.slug or f"provider-{provider.id}",
|
||||
provider_type=provider.provider_type,
|
||||
base_url=public_provider_url(provider.base_url),
|
||||
)
|
||||
for model in provider.models:
|
||||
key = (model.id.lower(), provider.id)
|
||||
if model.enabled:
|
||||
overrides_by_key[key] = model
|
||||
else:
|
||||
disabled_model_keys.add(key)
|
||||
|
||||
return overrides_by_key, disabled_model_keys, provider_identities
|
||||
|
||||
|
||||
def _apply_model_visibility(
|
||||
upstream: BaseUpstreamProvider,
|
||||
overrides_by_key: dict[ModelKey, ModelRow] | None,
|
||||
disabled_model_keys: set[ModelKey] | None,
|
||||
) -> list[object]:
|
||||
"""Return provider models after DB disabled/override state is applied.
|
||||
|
||||
Only the identity fields (``id``, ``forwarded_model_id``,
|
||||
``canonical_slug``) matter for path discovery, so DB override rows are used
|
||||
directly rather than rebuilt into fully priced ``Model`` objects — the
|
||||
pricing pipeline costs ~0.7ms of event-loop CPU per row for data this
|
||||
module immediately discards.
|
||||
"""
|
||||
overrides_by_key = overrides_by_key or {}
|
||||
disabled_model_keys = disabled_model_keys or set()
|
||||
upstream_provider_id = getattr(upstream, "db_id", None)
|
||||
if not isinstance(upstream_provider_id, int):
|
||||
return [
|
||||
model
|
||||
for model in upstream.get_cached_models()
|
||||
if getattr(model, "enabled", True)
|
||||
]
|
||||
|
||||
visible_models: list[object] = []
|
||||
seen_model_ids: set[str] = set()
|
||||
|
||||
for model in upstream.get_cached_models():
|
||||
model_id = getattr(model, "id", "")
|
||||
key = (model_id.lower(), upstream_provider_id)
|
||||
if not getattr(model, "enabled", True) or key in disabled_model_keys:
|
||||
continue
|
||||
# Apply overrides only for this provider's own model row.
|
||||
override_row = overrides_by_key.get(key)
|
||||
visible: object = model if override_row is None else override_row
|
||||
visible_models.append(visible)
|
||||
seen_model_ids.add(model_id.lower())
|
||||
|
||||
# DB-only override rows for this provider with no cached counterpart.
|
||||
for (model_id_lower, provider_id), override_row in overrides_by_key.items():
|
||||
if provider_id != upstream_provider_id:
|
||||
continue
|
||||
if model_id_lower in seen_model_ids:
|
||||
continue
|
||||
visible_models.append(override_row)
|
||||
seen_model_ids.add(model_id_lower)
|
||||
|
||||
return visible_models
|
||||
|
||||
|
||||
async def _collect_provider_paths(
|
||||
upstream: BaseUpstreamProvider,
|
||||
provider_identity: ConfiguredProviderIdentity,
|
||||
overrides_by_key: dict[ModelKey, ModelRow] | None = None,
|
||||
disabled_model_keys: set[ModelKey] | None = None,
|
||||
cycle: _RefreshCycleState | None = None,
|
||||
) -> ProviderPathSnapshot:
|
||||
"""Collect selectable routes while marking model-level degraded fetches.
|
||||
|
||||
A failed OpenRouter lookup preserves only that model's prior rows. Other
|
||||
models in the same provider still refresh, so a partial outage cannot erase
|
||||
valid discovery data or freeze the entire provider snapshot.
|
||||
"""
|
||||
cycle = cycle or _RefreshCycleState()
|
||||
models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys)
|
||||
|
||||
def _base_path(model: object) -> DiscoveredPath:
|
||||
model_id = exposed_model_id(model)
|
||||
return DiscoveredPath(
|
||||
model_id=model_id,
|
||||
path=encode_model_path(
|
||||
provider_identity.base_url, provider_identity.id, model_id
|
||||
),
|
||||
provider=provider_identity,
|
||||
)
|
||||
|
||||
if not is_openrouter_base_url(upstream.base_url):
|
||||
return ProviderPathSnapshot(paths=tuple(_base_path(model) for model in models))
|
||||
|
||||
if not (upstream.provider_type or "").strip():
|
||||
return ProviderPathSnapshot(paths=())
|
||||
|
||||
semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY)
|
||||
async with _make_http_client() as client:
|
||||
|
||||
async def _for_model(
|
||||
model: object,
|
||||
) -> tuple[list[DiscoveredPath], str | None]:
|
||||
model_id = exposed_model_id(model)
|
||||
author_slug = openrouter_author_slug(model)
|
||||
if not author_slug:
|
||||
return [_base_path(model)], None
|
||||
endpoints = await _fetch_openrouter_endpoint_subproviders(
|
||||
client,
|
||||
upstream.base_url,
|
||||
upstream.api_key,
|
||||
author_slug,
|
||||
semaphore,
|
||||
cycle,
|
||||
)
|
||||
if endpoints is None:
|
||||
return [], model_id
|
||||
paths = [_base_path(model)]
|
||||
paths.extend(
|
||||
DiscoveredPath(
|
||||
model_id=model_id,
|
||||
path=encode_model_path(
|
||||
provider_identity.base_url,
|
||||
provider_identity.id,
|
||||
model_id,
|
||||
endpoint.tag,
|
||||
),
|
||||
provider=provider_identity,
|
||||
endpoint_tag=endpoint.tag,
|
||||
endpoint_name=endpoint.provider_name,
|
||||
)
|
||||
for endpoint in endpoints
|
||||
)
|
||||
return paths, None
|
||||
|
||||
results = await asyncio.gather(
|
||||
*(_for_model(model) for model in models), return_exceptions=True
|
||||
)
|
||||
|
||||
paths: list[DiscoveredPath] = []
|
||||
preserve_model_ids: set[str] = set()
|
||||
for model, result in zip(models, results):
|
||||
if isinstance(result, BaseException):
|
||||
model_id = exposed_model_id(model)
|
||||
preserve_model_ids.add(model_id)
|
||||
logger.warning(
|
||||
"OpenRouter endpoint discovery task errored",
|
||||
extra={"provider": upstream.provider_type, "error": str(result)},
|
||||
)
|
||||
continue
|
||||
model_paths, preserved_model_id = result
|
||||
paths.extend(model_paths)
|
||||
if preserved_model_id:
|
||||
preserve_model_ids.add(preserved_model_id)
|
||||
|
||||
return ProviderPathSnapshot(
|
||||
paths=tuple(paths), preserve_model_ids=frozenset(preserve_model_ids)
|
||||
)
|
||||
|
||||
|
||||
async def _persist_provider_paths(
|
||||
upstream_provider_id: int, snapshot: ProviderPathSnapshot
|
||||
) -> None:
|
||||
"""Replace refreshed rows while retaining model-level degraded snapshots."""
|
||||
unique_paths = list(
|
||||
{(path.model_id, path.path): path for path in snapshot.paths}.values()
|
||||
)
|
||||
now = int(time.time())
|
||||
async with create_session() as session:
|
||||
delete_stmt = delete(ModelPathRow).where(
|
||||
col(ModelPathRow.upstream_provider_id) == upstream_provider_id
|
||||
)
|
||||
if snapshot.preserve_model_ids:
|
||||
delete_stmt = delete_stmt.where(
|
||||
col(ModelPathRow.model_id).not_in(sorted(snapshot.preserve_model_ids))
|
||||
)
|
||||
await session.exec(delete_stmt) # type: ignore[call-overload]
|
||||
for start in range(0, len(unique_paths), _PERSIST_CHUNK_SIZE):
|
||||
chunk = unique_paths[start : start + _PERSIST_CHUNK_SIZE]
|
||||
values = [
|
||||
{
|
||||
"model_id": discovered.model_id,
|
||||
"path": discovered.path,
|
||||
"provider_slug": discovered.provider.slug,
|
||||
"provider_type": discovered.provider.provider_type,
|
||||
"endpoint_tag": discovered.endpoint_tag,
|
||||
"endpoint_name": discovered.endpoint_name,
|
||||
"upstream_provider_id": upstream_provider_id,
|
||||
"updated_at": now,
|
||||
}
|
||||
for discovered in chunk
|
||||
]
|
||||
insert_stmt = insert(ModelPathRow).values(values)
|
||||
await session.execute(
|
||||
insert_stmt.on_conflict_do_update(
|
||||
index_elements=["model_id", "path", "upstream_provider_id"],
|
||||
set_={
|
||||
"provider_slug": insert_stmt.excluded.provider_slug,
|
||||
"provider_type": insert_stmt.excluded.provider_type,
|
||||
"endpoint_tag": insert_stmt.excluded.endpoint_tag,
|
||||
"endpoint_name": insert_stmt.excluded.endpoint_name,
|
||||
"updated_at": insert_stmt.excluded.updated_at,
|
||||
},
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def prune_model_paths_for_inactive_providers() -> None:
|
||||
"""Delete paths whose provider is no longer enabled in the database.
|
||||
|
||||
Called from ``refresh_model_maps`` so admin mutations (disable/delete
|
||||
provider) stop advertising a provider's paths immediately instead of
|
||||
waiting for the next timed refresh. Uses the DB as the source of truth, so
|
||||
it is safe at boot even before upstreams initialize.
|
||||
"""
|
||||
async with create_session() as session:
|
||||
enabled_ids = (
|
||||
await session.exec(
|
||||
select(UpstreamProviderRow.id).where(
|
||||
col(UpstreamProviderRow.enabled).is_(True)
|
||||
)
|
||||
)
|
||||
).all()
|
||||
stmt = delete(ModelPathRow)
|
||||
if enabled_ids:
|
||||
stmt = stmt.where(
|
||||
col(ModelPathRow.upstream_provider_id).not_in(
|
||||
[pid for pid in enabled_ids if pid is not None]
|
||||
)
|
||||
)
|
||||
await session.exec(stmt) # type: ignore[call-overload]
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def refresh_model_paths(
|
||||
upstreams: list[BaseUpstreamProvider],
|
||||
) -> None:
|
||||
"""Recompute and persist model paths for every enabled provider.
|
||||
|
||||
One provider's failure is logged and isolated; it must not break the rest.
|
||||
A provider whose paths could not be determined this cycle keeps its
|
||||
previously persisted rows. An empty ``upstreams`` list (e.g. a failed
|
||||
``initialize_upstreams`` at boot) is treated as "unknown" and touches
|
||||
nothing.
|
||||
"""
|
||||
if not upstreams:
|
||||
logger.warning("Skipping model paths refresh: no live upstreams")
|
||||
return
|
||||
|
||||
(
|
||||
overrides_by_key,
|
||||
disabled_model_keys,
|
||||
provider_identities,
|
||||
) = await _load_model_visibility()
|
||||
await prune_model_paths_for_inactive_providers()
|
||||
|
||||
cycle = _RefreshCycleState()
|
||||
for upstream in upstreams:
|
||||
if upstream.db_id is None or upstream.db_id not in provider_identities:
|
||||
continue
|
||||
try:
|
||||
snapshot = await _collect_provider_paths(
|
||||
upstream,
|
||||
provider_identity=provider_identities[upstream.db_id],
|
||||
overrides_by_key=overrides_by_key,
|
||||
disabled_model_keys=disabled_model_keys,
|
||||
cycle=cycle,
|
||||
)
|
||||
if snapshot.preserve_model_ids:
|
||||
logger.warning(
|
||||
"Some model paths are unknown; keeping their previous rows",
|
||||
extra={
|
||||
"provider": upstream.provider_type or upstream.base_url,
|
||||
"db_id": upstream.db_id,
|
||||
"preserved_models": len(snapshot.preserve_model_ids),
|
||||
},
|
||||
)
|
||||
await _persist_provider_paths(upstream.db_id, snapshot)
|
||||
except Exception as e: # noqa: BLE001 - isolate per-provider failures
|
||||
logger.error(
|
||||
"Failed to refresh model paths for provider",
|
||||
extra={
|
||||
"provider": upstream.provider_type or upstream.base_url,
|
||||
"db_id": upstream.db_id,
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None:
|
||||
"""Synchronize one provider when model-path discovery is enabled."""
|
||||
if _refresh_interval_seconds() <= 0:
|
||||
return
|
||||
|
||||
from ..proxy import get_upstreams
|
||||
|
||||
matching = [
|
||||
upstream
|
||||
for upstream in get_upstreams()
|
||||
if upstream.db_id == upstream_provider_id
|
||||
]
|
||||
if matching:
|
||||
await refresh_model_paths(matching)
|
||||
else:
|
||||
await prune_model_paths_for_inactive_providers()
|
||||
|
||||
|
||||
async def _drain_scheduled_provider_refreshes() -> None:
|
||||
"""Serialize and coalesce model-path refreshes scheduled by admin writes."""
|
||||
global _scheduled_provider_refresh_task
|
||||
|
||||
try:
|
||||
# Let mutations in the same event-loop turn collapse into one refresh.
|
||||
await asyncio.sleep(0)
|
||||
while _scheduled_provider_refresh_ids:
|
||||
if _refresh_interval_seconds() <= 0:
|
||||
_scheduled_provider_refresh_ids.clear()
|
||||
return
|
||||
provider_id = min(_scheduled_provider_refresh_ids)
|
||||
_scheduled_provider_refresh_ids.remove(provider_id)
|
||||
try:
|
||||
await refresh_model_paths_for_provider(provider_id)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001 - background best effort
|
||||
logger.warning(
|
||||
"Failed to refresh model paths after admin mutation",
|
||||
extra={
|
||||
"upstream_provider_id": provider_id,
|
||||
"error": str(exc),
|
||||
"error_type": type(exc).__name__,
|
||||
},
|
||||
)
|
||||
finally:
|
||||
_scheduled_provider_refresh_task = None
|
||||
|
||||
|
||||
async def schedule_model_paths_refresh_for_provider(
|
||||
upstream_provider_id: int,
|
||||
) -> None:
|
||||
"""Queue a non-blocking, coalesced refresh after an admin mutation."""
|
||||
global _scheduled_provider_refresh_task
|
||||
|
||||
if _refresh_interval_seconds() <= 0:
|
||||
return
|
||||
_scheduled_provider_refresh_ids.add(upstream_provider_id)
|
||||
if (
|
||||
_scheduled_provider_refresh_task is None
|
||||
or _scheduled_provider_refresh_task.done()
|
||||
):
|
||||
_scheduled_provider_refresh_task = asyncio.create_task(
|
||||
_drain_scheduled_provider_refreshes(),
|
||||
name="model-path-admin-refresh",
|
||||
)
|
||||
|
||||
|
||||
def _refresh_interval_seconds() -> int:
|
||||
"""Current interval, re-read every loop so runtime setting changes apply."""
|
||||
from ..core.settings import settings
|
||||
|
||||
if not getattr(settings, "enable_model_paths_refresh", True):
|
||||
return 0
|
||||
return int(getattr(settings, "model_paths_refresh_interval_seconds", 0) or 0)
|
||||
|
||||
|
||||
async def refresh_model_paths_periodically(
|
||||
upstreams_provider: (
|
||||
Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider]
|
||||
),
|
||||
) -> None:
|
||||
"""Background task mirroring ``refresh_upstreams_models_periodically``.
|
||||
|
||||
The interval and enable flag are re-read every iteration, so the refresh
|
||||
can be turned off (or on) and retuned without a restart. While disabled the
|
||||
task idles instead of exiting, so re-enabling takes effect.
|
||||
"""
|
||||
_DISABLED_POLL_SECONDS = 60.0
|
||||
|
||||
def _resolve_upstreams() -> list[BaseUpstreamProvider]:
|
||||
if callable(upstreams_provider):
|
||||
return upstreams_provider()
|
||||
return upstreams_provider
|
||||
|
||||
while True:
|
||||
interval = _refresh_interval_seconds()
|
||||
if interval <= 0:
|
||||
try:
|
||||
await asyncio.sleep(_DISABLED_POLL_SECONDS)
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
continue
|
||||
|
||||
try:
|
||||
await refresh_model_paths(_resolve_upstreams())
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.error(
|
||||
"Error in model paths refresh loop",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
try:
|
||||
jitter = max(0.0, float(interval) * 0.1)
|
||||
await asyncio.sleep(interval + random.uniform(0, jitter))
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
|
||||
|
||||
def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
|
||||
endpoint = None
|
||||
if row.endpoint_tag or row.endpoint_name:
|
||||
endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name}
|
||||
return {
|
||||
"path": row.path,
|
||||
"provider": {
|
||||
"id": row.upstream_provider_id,
|
||||
"slug": row.provider_slug,
|
||||
"type": row.provider_type,
|
||||
},
|
||||
"endpoint": endpoint,
|
||||
}
|
||||
|
||||
|
||||
async def get_all_model_paths() -> dict:
|
||||
"""All models with their exact selectable routes."""
|
||||
async with create_session() as session:
|
||||
rows = (
|
||||
await session.exec(
|
||||
select(ModelPathRow).order_by(
|
||||
col(ModelPathRow.model_id),
|
||||
col(ModelPathRow.path),
|
||||
col(ModelPathRow.upstream_provider_id),
|
||||
)
|
||||
)
|
||||
).all()
|
||||
|
||||
grouped: dict[str, list[dict[str, Any]]] = {}
|
||||
seen_paths: dict[str, set[str]] = {}
|
||||
updated_at = 0
|
||||
for row in rows:
|
||||
updated_at = max(updated_at, row.updated_at)
|
||||
if row.path in seen_paths.setdefault(row.model_id, set()):
|
||||
continue
|
||||
seen_paths[row.model_id].add(row.path)
|
||||
grouped.setdefault(row.model_id, []).append(_serialize_path(row))
|
||||
data = [
|
||||
{
|
||||
"id": grouped_model_id,
|
||||
"paths": grouped[grouped_model_id],
|
||||
}
|
||||
for grouped_model_id in sorted(grouped)
|
||||
]
|
||||
return {"data": data, "updated_at": updated_at or None}
|
||||
|
||||
|
||||
async def get_paths_for_model(model_id: str) -> dict:
|
||||
"""Return paths for an advertised ID or its provider-prefixed alias."""
|
||||
|
||||
async def load_rows(session: AsyncSession, lookup_id: str) -> list[ModelPathRow]:
|
||||
return list(
|
||||
(
|
||||
await session.exec(
|
||||
select(ModelPathRow)
|
||||
.where(col(ModelPathRow.model_id) == lookup_id)
|
||||
.order_by(
|
||||
col(ModelPathRow.path),
|
||||
col(ModelPathRow.upstream_provider_id),
|
||||
)
|
||||
)
|
||||
).all()
|
||||
)
|
||||
|
||||
async with create_session() as session:
|
||||
rows = await load_rows(session, model_id)
|
||||
if not rows:
|
||||
unprefixed_id = public_model_id(model_id)
|
||||
if unprefixed_id != model_id:
|
||||
rows = await load_rows(session, unprefixed_id)
|
||||
|
||||
seen: set[str] = set()
|
||||
paths: list[dict] = []
|
||||
updated_at = 0
|
||||
for row in rows:
|
||||
updated_at = max(updated_at, row.updated_at)
|
||||
if row.path in seen:
|
||||
continue
|
||||
seen.add(row.path)
|
||||
paths.append(_serialize_path(row))
|
||||
return {"data": paths, "updated_at": updated_at or None}
|
||||
@@ -1,4 +1,3 @@
|
||||
import json
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import httpx
|
||||
@@ -19,38 +18,6 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
|
||||
supports_anthropic_messages = True
|
||||
litellm_provider_prefix = "openrouter/"
|
||||
|
||||
def prepare_request_body(
|
||||
self, body: bytes | None, model_obj: Model
|
||||
) -> bytes | None:
|
||||
"""Set provider.require_parameters on tool-use requests.
|
||||
|
||||
Without it OpenRouter can route a tool call to an endpoint that doesn't
|
||||
support function calling and 404 with "No endpoints found that support
|
||||
tool use". We leave a client-supplied value untouched.
|
||||
"""
|
||||
body = super().prepare_request_body(body, model_obj)
|
||||
if not body:
|
||||
return body
|
||||
|
||||
try:
|
||||
data = json.loads(body)
|
||||
except json.JSONDecodeError:
|
||||
return body
|
||||
|
||||
if not isinstance(data, dict) or not data.get("tools"):
|
||||
return body
|
||||
|
||||
provider = data.get("provider")
|
||||
if not isinstance(provider, dict):
|
||||
provider = {}
|
||||
|
||||
if "require_parameters" in provider:
|
||||
return body
|
||||
|
||||
provider["require_parameters"] = True
|
||||
data["provider"] = provider
|
||||
return json.dumps(data).encode()
|
||||
|
||||
def _apply_provider_field(self, response_json: object) -> None:
|
||||
"""Stamp the ``provider`` field for OpenRouter responses.
|
||||
|
||||
|
||||
@@ -441,7 +441,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
"""
|
||||
data = await self.check_balance()
|
||||
balance = data.get("balance")
|
||||
if isinstance(balance, (int, float)):
|
||||
if isinstance(balance, (int, float)) and not isinstance(balance, bool):
|
||||
return float(balance)
|
||||
return None
|
||||
|
||||
|
||||
@@ -0,0 +1,203 @@
|
||||
"""Shared price/metadata resolution chain for upstream model discovery.
|
||||
|
||||
Most OpenAI-compatible ``/models`` responses carry no pricing. Rather than let
|
||||
a provider fabricate one, this module resolves a model through decreasingly
|
||||
trustworthy sources — litellm's bundled cost map (curated list prices, mirrors
|
||||
provider docs), then the OpenRouter feed (resale prices, broader coverage) —
|
||||
and returns ``None`` when none of them know the model, so the caller can fail
|
||||
closed instead of inventing a number.
|
||||
|
||||
Provider-native pricing (a gateway's own ``/models`` schema, e.g. Venice's
|
||||
``model_spec``) is authoritative and handled by the provider before this chain
|
||||
is consulted; only the shared fallback lives here so a later refactor can hoist
|
||||
it into the base provider unchanged.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResolvedPricing:
|
||||
"""Per-token pricing plus whatever metadata the answering source carried.
|
||||
|
||||
Prices are USD per token. ``source`` records provenance
|
||||
(``native``/``litellm``/``openrouter``/``unresolved``) so later work can
|
||||
surface where each price came from.
|
||||
"""
|
||||
|
||||
prompt: float
|
||||
completion: float
|
||||
context_length: int | None
|
||||
source: str
|
||||
modality: str | None = None
|
||||
max_completion_tokens: int | None = None
|
||||
input_cache_read: float = 0.0
|
||||
input_cache_write: float = 0.0
|
||||
input_modalities: list[str] = field(default_factory=lambda: ["text"])
|
||||
output_modalities: list[str] = field(default_factory=lambda: ["text"])
|
||||
tokenizer: str = "unknown"
|
||||
instruct_type: str | None = None
|
||||
is_moderated: bool | None = None
|
||||
|
||||
|
||||
def estimate_context_length(model_id: str) -> int:
|
||||
"""Best-effort context window from a model id when no source reports one.
|
||||
|
||||
The last rung of the fallback chain, reached only for a model whose price
|
||||
resolved but whose context did not (or that imported disabled). Context is
|
||||
not a billing input, so a rough id-based guess is acceptable here where a
|
||||
guessed *price* never would be.
|
||||
"""
|
||||
lowered = model_id.lower()
|
||||
if any(pattern in lowered for pattern in ["32k", "32000"]):
|
||||
return 32768
|
||||
if any(pattern in lowered for pattern in ["16k", "16000"]):
|
||||
return 16384
|
||||
if any(pattern in lowered for pattern in ["8k", "8000"]):
|
||||
return 8192
|
||||
if "gpt-4" in lowered:
|
||||
return 8192
|
||||
if "claude" in lowered:
|
||||
return 200000
|
||||
return 4096
|
||||
|
||||
|
||||
def _as_float(value: object) -> float | None:
|
||||
"""OpenRouter reports prices as strings; coerce, ``None`` if unparseable."""
|
||||
try:
|
||||
return float(value) # type: ignore[arg-type]
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _as_int(value: object) -> int | None:
|
||||
"""Coerce an already-numeric token count to ``int``, else ``None``."""
|
||||
return int(value) if isinstance(value, (int, float)) else None
|
||||
|
||||
|
||||
def _from_litellm(model_id: str) -> ResolvedPricing | None:
|
||||
# Lazy import so the resolver stays import-light and shares the exact
|
||||
# lookup semantics used by cache-rate backfill.
|
||||
from ..payment.models import litellm_cost_entry
|
||||
|
||||
info = litellm_cost_entry(model_id)
|
||||
if info is None:
|
||||
return None
|
||||
|
||||
prompt = info.get("input_cost_per_token")
|
||||
completion = info.get("output_cost_per_token")
|
||||
if not isinstance(prompt, (int, float)) or not isinstance(completion, (int, float)):
|
||||
return None
|
||||
# A both-zero entry is litellm listing a model without a real price (free
|
||||
# moderation/rerank tiers do this) — treating 0/0 as resolved would serve
|
||||
# the model for free. Reject it (and any negative) so the caller falls
|
||||
# through, mirroring async_fetch_openrouter_models' _has_valid_pricing.
|
||||
if prompt < 0 or completion < 0 or (prompt == 0 and completion == 0):
|
||||
return None
|
||||
|
||||
input_modalities = ["text"]
|
||||
if info.get("supports_vision"):
|
||||
input_modalities.append("image")
|
||||
|
||||
return ResolvedPricing(
|
||||
prompt=float(prompt),
|
||||
completion=float(completion),
|
||||
# max_input_tokens is the context window; max_tokens is litellm's
|
||||
# completion cap (it tracks max_output_tokens for ~94% of models), so
|
||||
# it is never a context source. A missing window falls to the id-based
|
||||
# estimate downstream rather than borrowing the output cap.
|
||||
context_length=_as_int(info.get("max_input_tokens")),
|
||||
source="litellm",
|
||||
max_completion_tokens=_as_int(info.get("max_output_tokens")),
|
||||
input_cache_read=float(info.get("cache_read_input_token_cost") or 0.0),
|
||||
input_cache_write=float(info.get("cache_creation_input_token_cost") or 0.0),
|
||||
input_modalities=input_modalities,
|
||||
)
|
||||
|
||||
|
||||
def _match_openrouter(model_id: str, feed: list[dict]) -> dict | None:
|
||||
"""Find ``model_id`` in the OpenRouter feed, exact id before bare tail.
|
||||
|
||||
Bare-tail matching (``deepseek-chat`` ↔ ``deepseek/deepseek-chat``) is a
|
||||
looser, lower-trust match — OpenRouter fans a model out across resellers —
|
||||
so an exact id match always wins first. When several entries share the bare
|
||||
tail, the one with the highest *combined* (prompt + completion) per-token
|
||||
cost wins: the choice must be deterministic (not feed-order-dependent) and
|
||||
money-safe whichever way traffic leans, since undercharging is the hazard.
|
||||
Ranking on prompt alone could pick an entry that is cheap on input but dear
|
||||
on output. The live feed has no such collisions today; this only governs
|
||||
the latent case.
|
||||
"""
|
||||
bare = model_id.split("/", 1)[-1]
|
||||
exact = next((m for m in feed if m.get("id") == model_id), None)
|
||||
if exact is not None:
|
||||
return exact
|
||||
matches = [m for m in feed if m.get("id", "").split("/", 1)[-1] == bare]
|
||||
if not matches:
|
||||
return None
|
||||
|
||||
def _combined_cost(m: dict) -> float:
|
||||
pricing = m.get("pricing", {})
|
||||
return (_as_float(pricing.get("prompt")) or 0.0) + (
|
||||
_as_float(pricing.get("completion")) or 0.0
|
||||
)
|
||||
|
||||
return max(matches, key=_combined_cost)
|
||||
|
||||
|
||||
def _from_openrouter(model_id: str, feed: list[dict]) -> ResolvedPricing | None:
|
||||
entry = _match_openrouter(model_id, feed)
|
||||
if entry is None:
|
||||
return None
|
||||
|
||||
pricing = entry.get("pricing", {})
|
||||
prompt = _as_float(pricing.get("prompt"))
|
||||
completion = _as_float(pricing.get("completion"))
|
||||
if prompt is None or completion is None:
|
||||
return None
|
||||
|
||||
architecture = entry.get("architecture", {})
|
||||
top_provider = entry.get("top_provider", {})
|
||||
|
||||
return ResolvedPricing(
|
||||
prompt=prompt,
|
||||
completion=completion,
|
||||
context_length=_as_int(entry.get("context_length")),
|
||||
source="openrouter",
|
||||
modality=architecture.get("modality"),
|
||||
max_completion_tokens=_as_int(top_provider.get("max_completion_tokens")),
|
||||
input_cache_read=_as_float(pricing.get("input_cache_read")) or 0.0,
|
||||
input_cache_write=_as_float(pricing.get("input_cache_write")) or 0.0,
|
||||
input_modalities=architecture.get("input_modalities") or ["text"],
|
||||
output_modalities=architecture.get("output_modalities") or ["text"],
|
||||
tokenizer=architecture.get("tokenizer") or "unknown",
|
||||
instruct_type=architecture.get("instruct_type"),
|
||||
is_moderated=top_provider.get("is_moderated"),
|
||||
)
|
||||
|
||||
|
||||
class FallbackPricingResolver:
|
||||
"""Resolves models via litellm → OpenRouter for one discovery pass.
|
||||
|
||||
The OpenRouter catalog is fetched at most once and only when a model
|
||||
actually misses litellm, so a provider full of litellm-known models never
|
||||
touches the network. Instantiate one per ``fetch_models`` call.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._openrouter_feed: list[dict] | None = None
|
||||
|
||||
async def resolve(self, model_id: str) -> ResolvedPricing | None:
|
||||
"""Resolve ``model_id``; ``None`` if no source knows it."""
|
||||
resolved = _from_litellm(model_id)
|
||||
if resolved is not None:
|
||||
return resolved
|
||||
|
||||
if self._openrouter_feed is None:
|
||||
# Lazy import so tests can patch the feed at its source.
|
||||
from ..payment.models import async_fetch_openrouter_models
|
||||
|
||||
self._openrouter_feed = await async_fetch_openrouter_models()
|
||||
return _from_openrouter(model_id, self._openrouter_feed)
|
||||
@@ -74,7 +74,7 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
def _build_from_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "TinfoilUpstreamProvider":
|
||||
return cls(
|
||||
@@ -116,7 +116,7 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider):
|
||||
EHBP-only header used for encrypted POST requests and is not honored
|
||||
for unencrypted GET requests.
|
||||
"""
|
||||
clean_path = path.removeprefix("tee/")
|
||||
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)
|
||||
|
||||
@@ -27,6 +27,16 @@ _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
|
||||
@@ -47,6 +57,19 @@ def _get_header(headers: list[tuple[str, str]], name: str) -> str | None:
|
||||
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,
|
||||
@@ -71,6 +94,11 @@ async def forward_with_trailer(
|
||||
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),
|
||||
@@ -86,7 +114,7 @@ async def forward_with_trailer(
|
||||
header_lines.append("Connection: close")
|
||||
|
||||
for key, value in headers.items():
|
||||
if key.lower() in ("host", "connection"):
|
||||
if key.lower() == "host":
|
||||
continue
|
||||
header_lines.append(f"{key}: {value}")
|
||||
|
||||
|
||||
+1870
-308
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
@@ -203,8 +203,13 @@ class TestmintWallet:
|
||||
token_base64 = base64.urlsafe_b64encode(token_json.encode()).decode()
|
||||
return f"cashuA{token_base64}"
|
||||
|
||||
async def redeem_token(self, token: str) -> Tuple[int, str, str]:
|
||||
"""Redeem a Cashu token - compatible with wallet.recieve_token"""
|
||||
async def redeem_token(
|
||||
self,
|
||||
token: str,
|
||||
destination_mint: str | None = None,
|
||||
destination_unit: str | None = None,
|
||||
) -> Tuple[int, str, str]:
|
||||
"""Redeem a Cashu token - compatible with wallet.recieve_token."""
|
||||
if not self.wallet:
|
||||
await self.init()
|
||||
|
||||
|
||||
@@ -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"
|
||||
@@ -15,6 +15,7 @@ from unittest.mock import patch
|
||||
import pytest
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.auth import ReservationSnapshot
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.payment.cost_calculation import CostData
|
||||
|
||||
@@ -23,7 +24,7 @@ def _make_key(balance: int, reserved: int) -> ApiKey:
|
||||
return ApiKey(
|
||||
hashed_key=f"test_{uuid.uuid4().hex}",
|
||||
balance=balance,
|
||||
reserved_balance=reserved,
|
||||
reserved_balance=0,
|
||||
total_spent=0,
|
||||
total_requests=1,
|
||||
)
|
||||
@@ -75,6 +76,8 @@ async def test_balance_never_negative_when_cost_exceeds_reservation(
|
||||
key = _make_key(balance=deducted_max_cost, reserved=deducted_max_cost)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
from routstr.auth import pay_for_request
|
||||
await pay_for_request(key, deducted_max_cost, integration_session)
|
||||
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}}
|
||||
|
||||
@@ -82,7 +85,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)
|
||||
|
||||
@@ -109,6 +114,8 @@ async def test_balance_floor_at_zero_on_overrun(
|
||||
key = _make_key(balance=500, reserved=500)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
from routstr.auth import pay_for_request
|
||||
await pay_for_request(key, deducted_max_cost, integration_session)
|
||||
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
|
||||
|
||||
@@ -116,7 +123,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)
|
||||
|
||||
@@ -148,6 +157,8 @@ async def test_full_cost_charged_when_balance_sufficient_for_overrun(
|
||||
key = _make_key(balance=2000, reserved=990)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
from routstr.auth import pay_for_request
|
||||
await pay_for_request(key, deducted_max_cost, integration_session)
|
||||
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}}
|
||||
|
||||
@@ -155,7 +166,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)
|
||||
|
||||
@@ -184,7 +197,11 @@ async def test_concurrent_cost_overruns_never_negative(
|
||||
"""Concurrent finalization with cost overruns must never produce negative balance."""
|
||||
import asyncio
|
||||
|
||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
||||
from routstr.auth import (
|
||||
adjust_payment_for_tokens,
|
||||
get_reservation_snapshot,
|
||||
pay_for_request,
|
||||
)
|
||||
from routstr.core.db import create_session
|
||||
|
||||
deducted_max_cost = 990
|
||||
@@ -210,12 +227,14 @@ async def test_concurrent_cost_overruns_never_negative(
|
||||
async with create_session() as session:
|
||||
key_to_reserve = await session.get(ApiKey, key_hash)
|
||||
assert key_to_reserve is not None
|
||||
reservations = []
|
||||
for _ in range(n_requests):
|
||||
await pay_for_request(key_to_reserve, deducted_max_cost, session)
|
||||
reservations.append(await get_reservation_snapshot(key_to_reserve, session))
|
||||
await session.refresh(key_to_reserve)
|
||||
|
||||
# Now finalize all concurrently with cost overrun
|
||||
async def finalize() -> None:
|
||||
async def finalize(reservation: ReservationSnapshot) -> None:
|
||||
response_data = {
|
||||
"model": "test-model",
|
||||
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
|
||||
@@ -223,15 +242,22 @@ 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,
|
||||
reservation_snapshot=reservation,
|
||||
)
|
||||
|
||||
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(r) for r in reservations))
|
||||
|
||||
async with create_session() as session:
|
||||
final_key = await session.get(ApiKey, key_hash)
|
||||
@@ -272,6 +298,8 @@ async def test_zero_free_balance_overrun_is_safe(
|
||||
key = _make_key(balance=1000, reserved=1000)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
from routstr.auth import pay_for_request
|
||||
await pay_for_request(key, deducted_max_cost, integration_session)
|
||||
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 100}}
|
||||
|
||||
@@ -279,7 +307,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)
|
||||
|
||||
@@ -308,7 +338,11 @@ async def test_parallel_requests_no_free_inference(
|
||||
"""Second parallel finalization must be charged even when first depleted free balance."""
|
||||
import asyncio
|
||||
|
||||
from routstr.auth import adjust_payment_for_tokens
|
||||
from routstr.auth import (
|
||||
adjust_payment_for_tokens,
|
||||
get_reservation_snapshot,
|
||||
pay_for_request,
|
||||
)
|
||||
from routstr.core.db import create_session
|
||||
|
||||
deducted_max_cost = 100
|
||||
@@ -329,14 +363,18 @@ async def test_parallel_requests_no_free_inference(
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=starting_balance,
|
||||
reserved_balance=deducted_max_cost * 2, # both slots pre-reserved
|
||||
reserved_balance=0,
|
||||
total_spent=0,
|
||||
total_requests=2,
|
||||
)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
reservations = []
|
||||
for _ in range(2):
|
||||
await pay_for_request(key, deducted_max_cost, session)
|
||||
reservations.append(await get_reservation_snapshot(key, session))
|
||||
|
||||
async def finalize() -> None:
|
||||
async def finalize(reservation: ReservationSnapshot) -> None:
|
||||
response_data = {
|
||||
"model": "test-model",
|
||||
"usage": {"prompt_tokens": 50, "completion_tokens": 100},
|
||||
@@ -344,15 +382,22 @@ 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,
|
||||
reservation_snapshot=reservation,
|
||||
)
|
||||
|
||||
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(r) for r in reservations))
|
||||
|
||||
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,666 @@
|
||||
"""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 sqlmodel import select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey, ReservationRelease
|
||||
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"
|
||||
THIRD_BASE_URL = "https://third.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, max_cost=100.0),
|
||||
)
|
||||
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],
|
||||
integration_session: AsyncSession,
|
||||
) -> 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
|
||||
|
||||
# The fallback's larger max-cost envelope requires a replacement
|
||||
# reservation. The failed candidate is released, the serving candidate is
|
||||
# charged, and no request-owned reservation remains active.
|
||||
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
|
||||
records = (
|
||||
await integration_session.exec(
|
||||
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
|
||||
)
|
||||
).all()
|
||||
assert sorted(record.status for record in records) == ["charged", "released"]
|
||||
|
||||
|
||||
@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 three_candidate_child_maps(
|
||||
patched_db_engine: None,
|
||||
) -> AsyncGenerator[None, None]:
|
||||
"""Second candidate cannot fit the child limit; third restores and serves."""
|
||||
first = _StaticProvider(
|
||||
CHEAP_BASE_URL,
|
||||
"key-first",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.001, 0.002, max_cost=50.0),
|
||||
)
|
||||
too_large = _StaticProvider(
|
||||
EXPENSIVE_BASE_URL,
|
||||
"key-too-large",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.002, 0.003, max_cost=100.0),
|
||||
)
|
||||
third = _StaticProvider(
|
||||
THIRD_BASE_URL,
|
||||
"key-third",
|
||||
1.0,
|
||||
_make_model("dual-model", 0.003, 0.004, max_cost=50.0),
|
||||
)
|
||||
async for _ in _install_providers([first, too_large, third]):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_child_failover_rolls_back_failed_larger_reserve_before_restoring(
|
||||
authenticated_client: AsyncClient,
|
||||
three_candidate_child_maps: None,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""A failed child guard cannot leak its parent update into restoration."""
|
||||
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
|
||||
child = await integration_session.get(ApiKey, key_hash)
|
||||
assert child is not None
|
||||
parent = ApiKey(hashed_key="failover-parent", balance=10_000_000)
|
||||
child.parent_key_hash = parent.hashed_key
|
||||
child.balance_limit = 75_000
|
||||
integration_session.add(parent)
|
||||
integration_session.add(child)
|
||||
await integration_session.commit()
|
||||
|
||||
sent_requests: list[httpx.Request] = []
|
||||
|
||||
async def fake_transport(
|
||||
request: httpx.Request, *args: Any, **kwargs: Any
|
||||
) -> httpx.Response:
|
||||
sent_requests.append(request)
|
||||
return _upstream_response(request)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
side_effect=fake_transport,
|
||||
),
|
||||
patch(
|
||||
"routstr.payment.cost_calculation.sats_usd_price",
|
||||
return_value=0.0005,
|
||||
),
|
||||
):
|
||||
response = await authenticated_client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": "dual-model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
# The 100-sat candidate is rejected before forwarding; the third serves.
|
||||
assert [request.url.host for request in sent_requests] == [
|
||||
"cheap.example.com",
|
||||
"third.example.com",
|
||||
]
|
||||
|
||||
await integration_session.refresh(parent)
|
||||
await integration_session.refresh(child)
|
||||
assert parent.reserved_balance == 0
|
||||
assert child.reserved_balance == 0
|
||||
assert parent.total_spent == response.json()["cost"]["total_msats"]
|
||||
|
||||
records = (
|
||||
await integration_session.exec(
|
||||
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
|
||||
)
|
||||
).all()
|
||||
assert len(records) == 2
|
||||
assert sorted(record.status for record in records) == ["charged", "released"]
|
||||
assert len({record.reserved_msats for record in records}) == 1
|
||||
assert all(record.status != "active" for record in records)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def raised_envelope_provider_maps(
|
||||
patched_db_engine: None,
|
||||
) -> 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,
|
||||
integration_session: AsyncSession,
|
||||
) -> 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
|
||||
|
||||
key_hash = authenticated_client._test_api_key.removeprefix("sk-") # type: ignore[attr-defined]
|
||||
records = (
|
||||
await integration_session.exec(
|
||||
select(ReservationRelease).where(ReservationRelease.key_hash == key_hash)
|
||||
)
|
||||
).all()
|
||||
assert len(records) == 2
|
||||
released = next(record for record in records if record.status == "released")
|
||||
charged = next(record for record in records if record.status == "charged")
|
||||
assert charged.reserved_msats > released.reserved_msats
|
||||
assert all(record.status != "active" for record in records)
|
||||
@@ -38,7 +38,7 @@ async def test_overrun_charges_after_reservation_swept(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Overrun finalize must charge even when the reservation was already released."""
|
||||
from routstr.auth import adjust_payment_for_tokens
|
||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
||||
|
||||
deducted_max_cost = 990 # discounted reservation
|
||||
actual_token_cost = 1000 # actual cost overruns the reservation
|
||||
@@ -47,6 +47,10 @@ async def test_overrun_charges_after_reservation_swept(
|
||||
key = _make_key(balance=1000, reserved=0)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
await pay_for_request(key, deducted_max_cost, integration_session)
|
||||
key.reserved_balance = 0
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
response_data = {
|
||||
"model": "test-model",
|
||||
@@ -58,7 +62,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)
|
||||
@@ -79,8 +83,16 @@ async def test_free_response_path_closed_end_to_end(
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""A reservation released by the real sweeper must not yield a free response."""
|
||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
||||
from routstr.core.db import create_session, release_stale_reservations
|
||||
from routstr.auth import (
|
||||
adjust_payment_for_tokens,
|
||||
get_reservation_snapshot,
|
||||
pay_for_request,
|
||||
)
|
||||
from routstr.core.db import (
|
||||
ReservationRelease,
|
||||
create_session,
|
||||
release_stale_reservations,
|
||||
)
|
||||
|
||||
deducted_max_cost = 990
|
||||
actual_token_cost = 1000
|
||||
@@ -104,10 +116,15 @@ async def test_free_response_path_closed_end_to_end(
|
||||
key = await session.get(ApiKey, key_hash)
|
||||
assert key is not None
|
||||
await pay_for_request(key, deducted_max_cost, session)
|
||||
snapshot = await get_reservation_snapshot(key, session)
|
||||
await session.refresh(key)
|
||||
assert key.reserved_balance == deducted_max_cost
|
||||
key.reserved_at = int(time.time()) - 10_000
|
||||
record = await session.get(ReservationRelease, snapshot.release_id)
|
||||
assert record is not None
|
||||
record.created_at = int(time.time()) - 10_000
|
||||
session.add(key)
|
||||
session.add(record)
|
||||
await session.commit()
|
||||
|
||||
# Sweeper releases the stale reservation without charging.
|
||||
@@ -129,18 +146,20 @@ 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,
|
||||
reservation_snapshot=snapshot,
|
||||
)
|
||||
|
||||
async with create_session() as session:
|
||||
final = await session.get(ApiKey, key_hash)
|
||||
assert final is not None
|
||||
|
||||
assert final.total_spent == actual_token_cost, (
|
||||
f"Free response: total_spent={final.total_spent}, expected {actual_token_cost}"
|
||||
)
|
||||
assert final.balance == 1000 - actual_token_cost, (
|
||||
f"Balance not charged after sweep: {final.balance}"
|
||||
)
|
||||
# Stale release is terminal for this reservation. A late finalizer must not
|
||||
# charge aggregate balance that may now belong to a newer request.
|
||||
assert final.total_spent == 0
|
||||
assert final.balance == 1000
|
||||
assert final.balance >= 0
|
||||
assert final.reserved_balance == 0
|
||||
|
||||
@@ -207,8 +207,35 @@ async def test_pay_for_request_succeeds_when_balance_equals_cost(
|
||||
assert key.balance == model_cost # balance unchanged, only reserved goes up
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_model_maximum_is_required_and_reserved(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
from routstr.auth import pay_for_request, validate_bearer_key
|
||||
|
||||
short_key = _key(balance=95_000)
|
||||
exact_key = _key(balance=100_000)
|
||||
integration_session.add(short_key)
|
||||
integration_session.add(exact_key)
|
||||
await integration_session.commit()
|
||||
|
||||
with pytest.raises(HTTPException) as insufficient:
|
||||
await validate_bearer_key(
|
||||
f"sk-{short_key.hashed_key}", integration_session, min_cost=100_000
|
||||
)
|
||||
assert insufficient.value.status_code == 402
|
||||
|
||||
validated = await validate_bearer_key(
|
||||
f"sk-{exact_key.hashed_key}", integration_session, min_cost=100_000
|
||||
)
|
||||
await pay_for_request(validated, 100_000, integration_session)
|
||||
|
||||
await integration_session.refresh(exact_key)
|
||||
assert exact_key.reserved_balance == 100_000
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 6 — HTTP layer returns 402 JSON with the right shape
|
||||
# HTTP layer returns 402 JSON with the right shape
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -241,8 +268,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",
|
||||
@@ -264,8 +293,8 @@ async def test_http_402_response_shape_on_insufficient_balance(
|
||||
error = body["detail"]["error"]
|
||||
assert error["code"] == "insufficient_balance"
|
||||
assert error["type"] == "insufficient_quota"
|
||||
assert str(model_cost) in error["message"]
|
||||
assert str(user_balance) in error["message"]
|
||||
assert "622.888 sats (622888 msats) required" in error["message"]
|
||||
assert "20.32 sats (20320 msats) available" in error["message"]
|
||||
|
||||
# Balance must be completely untouched
|
||||
await integration_session.refresh(key)
|
||||
|
||||
@@ -3,20 +3,30 @@
|
||||
Covers two things:
|
||||
- The three constraint fields (balance_limit, balance_limit_reset, validity_date)
|
||||
are persisted on LightningInvoice and survive a DB round-trip.
|
||||
- create_api_key_from_invoice propagates those fields to the created ApiKey,
|
||||
so the constraints are actually enforced when the key is used.
|
||||
- The production-path API-key record helper propagates those fields to the
|
||||
created ApiKey, so the constraints are actually enforced when the key is used.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from cashu.core.base import Proof
|
||||
from sqlalchemy import inspect
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey, LightningInvoice
|
||||
from routstr.lightning import create_api_key_from_invoice
|
||||
from routstr.lightning import _create_api_key_record
|
||||
|
||||
|
||||
def _configure_quote_proof_wallet(wallet: MagicMock) -> None:
|
||||
wallet.proofs = []
|
||||
wallet.keysets = {}
|
||||
wallet.load_proofs = AsyncMock()
|
||||
|
||||
|
||||
def _make_invoice(**kwargs: object) -> LightningInvoice:
|
||||
@@ -39,7 +49,15 @@ def _make_invoice(**kwargs: object) -> LightningInvoice:
|
||||
def mock_wallet_mint() -> object:
|
||||
with patch("routstr.lightning.get_wallet") as mock_get_wallet:
|
||||
wallet = AsyncMock()
|
||||
wallet.mint = AsyncMock(return_value=[])
|
||||
wallet.proofs = []
|
||||
wallet.load_proofs = AsyncMock()
|
||||
|
||||
async def mint(amount: int, quote_id: str) -> list[Proof]:
|
||||
proofs = [Proof(amount=amount, mint_id=quote_id)]
|
||||
wallet.proofs.extend(proofs)
|
||||
return proofs
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=mint)
|
||||
mock_get_wallet.return_value = wallet
|
||||
yield mock_get_wallet
|
||||
|
||||
@@ -48,6 +66,7 @@ def mock_wallet_mint() -> object:
|
||||
# Persistence
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_persists_balance_limit(
|
||||
integration_session: AsyncSession,
|
||||
@@ -92,6 +111,7 @@ async def test_invoice_persists_validity_date(
|
||||
# Propagation to ApiKey
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_created_key_receives_balance_limit(
|
||||
integration_session: AsyncSession,
|
||||
@@ -100,7 +120,7 @@ async def test_created_key_receives_balance_limit(
|
||||
integration_session.add(invoice)
|
||||
await integration_session.flush()
|
||||
|
||||
api_key = await create_api_key_from_invoice(invoice, integration_session)
|
||||
api_key = await _create_api_key_record(invoice, integration_session)
|
||||
await integration_session.commit()
|
||||
|
||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||
@@ -116,7 +136,7 @@ async def test_created_key_receives_balance_limit_reset(
|
||||
integration_session.add(invoice)
|
||||
await integration_session.flush()
|
||||
|
||||
api_key = await create_api_key_from_invoice(invoice, integration_session)
|
||||
api_key = await _create_api_key_record(invoice, integration_session)
|
||||
await integration_session.commit()
|
||||
|
||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||
@@ -133,7 +153,7 @@ async def test_created_key_receives_validity_date(
|
||||
integration_session.add(invoice)
|
||||
await integration_session.flush()
|
||||
|
||||
api_key = await create_api_key_from_invoice(invoice, integration_session)
|
||||
api_key = await _create_api_key_record(invoice, integration_session)
|
||||
await integration_session.commit()
|
||||
|
||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||
@@ -141,6 +161,254 @@ async def test_created_key_receives_validity_date(
|
||||
assert stored_key.validity_date == expiry
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_payment_check_releases_connection_during_mint_quote(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(id="inv_slow_quote", status="pending", paid_at=None)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
|
||||
async def quote_status(*args: object, **kwargs: object) -> MagicMock:
|
||||
assert integration_engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
return MagicMock(paid=False)
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.get_mint_quote = AsyncMock(side_effect=quote_status)
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await check_invoice_payment(stored, session)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_payment_checks_mint_and_credit_invoice_once(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(id="inv_concurrent", status="pending", paid_at=None)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
_configure_quote_proof_wallet(wallet)
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
|
||||
mint_calls = 0
|
||||
|
||||
async def single_use_mint(*args: object, **kwargs: object) -> list[object]:
|
||||
# Real mints enforce single-use quotes: the second concurrent minter
|
||||
# gets rejected at the mint, mirroring cashu quote semantics.
|
||||
nonlocal mint_calls
|
||||
mint_calls += 1
|
||||
call_number = mint_calls
|
||||
await asyncio.sleep(0.05)
|
||||
if call_number > 1:
|
||||
raise Exception("quote already issued")
|
||||
proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash)
|
||||
wallet.proofs.append(proof)
|
||||
return [proof]
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=single_use_mint)
|
||||
|
||||
async with (
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as first,
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as second,
|
||||
):
|
||||
first_invoice = await first.get(LightningInvoice, invoice.id)
|
||||
second_invoice = await second.get(LightningInvoice, invoice.id)
|
||||
assert first_invoice is not None
|
||||
assert second_invoice is not None
|
||||
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await asyncio.gather(
|
||||
check_invoice_payment(first_invoice, first),
|
||||
check_invoice_payment(second_invoice, second),
|
||||
)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored_invoice = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored_invoice is not None
|
||||
assert stored_invoice.status == "paid"
|
||||
assert stored_invoice.api_key_hash is not None
|
||||
stored_key = await verify.get(ApiKey, stored_invoice.api_key_hash)
|
||||
assert stored_key is not None
|
||||
assert stored_key.balance == invoice.amount_sats * 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_mint_marks_invoice_for_settlement_retry(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(id="inv_mint_failure", status="pending", paid_at=None)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
_configure_quote_proof_wallet(wallet)
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
wallet.mint = AsyncMock(side_effect=TimeoutError("mint unavailable"))
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await check_invoice_payment(stored, session)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
assert stored.status == "settlement_pending"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unpaid_topup_does_not_query_target_key(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(
|
||||
id="inv_unpaid_topup",
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
purpose="topup",
|
||||
api_key_hash="target-key",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=False))
|
||||
create_session = MagicMock(side_effect=RuntimeError("target lookup should not run"))
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.lightning.create_session", create_session),
|
||||
):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await check_invoice_payment(stored, session)
|
||||
|
||||
wallet.get_mint_quote.assert_awaited_once_with(invoice.payment_hash)
|
||||
create_session.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_topup_target_is_rejected_before_mint(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(
|
||||
id="inv_missing_topup_target",
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
purpose="topup",
|
||||
api_key_hash="pruned-key",
|
||||
expires_at=int(time.time()) - 1,
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
wallet.mint = AsyncMock()
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.lightning.logger.critical") as critical,
|
||||
):
|
||||
from routstr.lightning import get_invoice_status
|
||||
|
||||
response = await get_invoice_status(invoice.id, session)
|
||||
|
||||
assert response.status == "reconciliation_required"
|
||||
assert stored.status == "reconciliation_required"
|
||||
assert stored not in session.dirty
|
||||
critical.assert_called_once()
|
||||
|
||||
wallet.mint.assert_not_awaited()
|
||||
async with AsyncSession(integration_engine) as verify:
|
||||
stored = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
assert stored.status == "reconciliation_required"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_mint_db_failure_keeps_invoice_pending_for_reconciliation(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(id="inv_finalize_failure", status="pending", paid_at=None)
|
||||
sibling = _make_invoice(
|
||||
id="inv_finalize_failure_sibling",
|
||||
bolt11="lnbc1000n1sibling",
|
||||
payment_hash="cafebabe" * 8,
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add_all([invoice, sibling])
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
_configure_quote_proof_wallet(wallet)
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
|
||||
async def successful_mint(*args: object, **kwargs: object) -> list[Proof]:
|
||||
proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash)
|
||||
wallet.proofs.append(proof)
|
||||
return [proof]
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=successful_mint)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
stored_sibling = await session.get(LightningInvoice, sibling.id)
|
||||
assert stored is not None
|
||||
assert stored_sibling is not None
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch(
|
||||
"routstr.lightning._create_api_key_record",
|
||||
AsyncMock(side_effect=RuntimeError("database unavailable")),
|
||||
),
|
||||
):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await check_invoice_payment(stored, session)
|
||||
|
||||
stored_state = inspect(stored)
|
||||
sibling_state = inspect(stored_sibling)
|
||||
assert stored_state is not None
|
||||
assert sibling_state is not None
|
||||
assert stored_state.expired is False
|
||||
assert sibling_state.expired is False
|
||||
assert stored.status == "settlement_pending"
|
||||
assert stored_sibling.id == sibling.id
|
||||
|
||||
assert wallet.mint.await_count == 1
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
assert stored.status == "settlement_pending"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_created_key_without_constraints_has_none_fields(
|
||||
integration_session: AsyncSession,
|
||||
@@ -149,7 +417,7 @@ async def test_created_key_without_constraints_has_none_fields(
|
||||
integration_session.add(invoice)
|
||||
await integration_session.flush()
|
||||
|
||||
api_key = await create_api_key_from_invoice(invoice, integration_session)
|
||||
api_key = await _create_api_key_record(invoice, integration_session)
|
||||
await integration_session.commit()
|
||||
|
||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||
@@ -157,3 +425,87 @@ async def test_created_key_without_constraints_has_none_fields(
|
||||
assert stored_key.balance_limit is None
|
||||
assert stored_key.balance_limit_reset is None
|
||||
assert stored_key.validity_date is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_guard_credits_once_when_both_mints_succeed(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Even if the mint fails to enforce single-use quotes and both racers
|
||||
mint successfully, the conditional status update must credit exactly once."""
|
||||
key = ApiKey(hashed_key="race-key", balance=1_000)
|
||||
invoice = _make_invoice(
|
||||
id="inv_db_guard",
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
purpose="topup",
|
||||
api_key_hash="race-key",
|
||||
)
|
||||
sibling = _make_invoice(
|
||||
id="inv_db_guard_sibling",
|
||||
bolt11="lnbc1000n1race-sibling",
|
||||
payment_hash="01234567" * 8,
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add_all([key, invoice, sibling])
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
_configure_quote_proof_wallet(wallet)
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
|
||||
async def always_succeeding_mint(*args: object, **kwargs: object) -> list[Proof]:
|
||||
await asyncio.sleep(0.05)
|
||||
proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash)
|
||||
wallet.proofs.append(proof)
|
||||
return [proof]
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=always_succeeding_mint)
|
||||
|
||||
async with (
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as first,
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as second,
|
||||
):
|
||||
first_invoice = await first.get(LightningInvoice, invoice.id)
|
||||
first_sibling = await first.get(LightningInvoice, sibling.id)
|
||||
second_invoice = await second.get(LightningInvoice, invoice.id)
|
||||
assert first_invoice is not None
|
||||
assert first_sibling is not None
|
||||
assert second_invoice is not None
|
||||
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await asyncio.gather(
|
||||
check_invoice_payment(first_invoice, first),
|
||||
check_invoice_payment(second_invoice, second),
|
||||
)
|
||||
|
||||
first_state = inspect(first_invoice)
|
||||
sibling_state = inspect(first_sibling)
|
||||
second_state = inspect(second_invoice)
|
||||
assert first_state is not None
|
||||
assert sibling_state is not None
|
||||
assert second_state is not None
|
||||
assert first_state.expired is False
|
||||
assert sibling_state.expired is False
|
||||
assert second_state.expired is False
|
||||
assert first_invoice.id == invoice.id
|
||||
assert first_sibling.id == sibling.id
|
||||
assert second_invoice.id == invoice.id
|
||||
assert first_invoice.status == "paid"
|
||||
assert second_invoice.status == "paid"
|
||||
assert first_invoice not in first.dirty
|
||||
assert second_invoice not in second.dirty
|
||||
|
||||
assert wallet.mint.await_count == 1
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored_invoice = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored_invoice is not None
|
||||
assert stored_invoice.status == "paid"
|
||||
stored_key = await verify.get(ApiKey, "race-key")
|
||||
assert stored_key is not None
|
||||
assert stored_key.balance == 1_000 + invoice.amount_sats * 1000
|
||||
|
||||
@@ -26,11 +26,17 @@ async def patch_invoice_generation() -> Any:
|
||||
"""Stub out `generate_lightning_invoice` so no mint round-trip is needed."""
|
||||
counter = {"n": 0}
|
||||
|
||||
async def fake_generate(amount_sats: int, description: str) -> tuple[str, str]:
|
||||
async def fake_generate(
|
||||
amount_sats: int,
|
||||
description: str,
|
||||
*,
|
||||
allowed_mints: list[str] | None = None,
|
||||
) -> tuple[str, str, str]:
|
||||
counter["n"] += 1
|
||||
return (
|
||||
f"lnbc{amount_sats}n1pfakeinvoice{counter['n']}",
|
||||
f"payment_hash_{counter['n']}",
|
||||
"http://localhost:3338",
|
||||
)
|
||||
|
||||
with patch(
|
||||
@@ -95,6 +101,8 @@ async def test_topup_with_authorization_header(
|
||||
body = resp.json()
|
||||
assert body["amount_sats"] == 500
|
||||
assert body["bolt11"].startswith("lnbc")
|
||||
allowed_mints = patch_invoice_generation.call_args.kwargs["allowed_mints"]
|
||||
assert allowed_mints == ["http://localhost:3338"]
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
|
||||
@@ -0,0 +1,365 @@
|
||||
import asyncio
|
||||
import time
|
||||
import uuid
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from cashu.core.base import Proof
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlmodel import col, update
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey, LightningInvoice
|
||||
from routstr.lightning import (
|
||||
_expire_invoice_if_authoritatively_unpaid,
|
||||
_finalize_invoice_settlement,
|
||||
_InvoiceSettlement,
|
||||
check_invoice_payment,
|
||||
)
|
||||
|
||||
|
||||
def _lightning_invoice(**overrides: object) -> LightningInvoice:
|
||||
suffix = uuid.uuid4().hex
|
||||
values = {
|
||||
"id": f"invoice-{suffix}",
|
||||
"bolt11": f"lnbc-{suffix}",
|
||||
"amount_sats": 100,
|
||||
"description": "settlement test",
|
||||
"payment_hash": f"quote-{suffix}",
|
||||
"status": "pending",
|
||||
"purpose": "create",
|
||||
"mint_url": "http://mint:3338",
|
||||
"expires_at": int(time.time()) + 3600,
|
||||
}
|
||||
values.update(overrides)
|
||||
return LightningInvoice(**values) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_read_transaction_closes_before_external_mint_io(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
invoice = _lightning_invoice()
|
||||
integration_session.add(invoice)
|
||||
await integration_session.commit()
|
||||
stored = await integration_session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
|
||||
wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=False)))
|
||||
|
||||
async def get_wallet_without_open_db_transaction(
|
||||
*args: object, **kwargs: object
|
||||
) -> Mock:
|
||||
assert not integration_session.in_transaction()
|
||||
return wallet
|
||||
|
||||
with patch(
|
||||
"routstr.lightning.get_wallet", side_effect=get_wallet_without_open_db_transaction
|
||||
):
|
||||
await check_invoice_payment(stored, integration_session)
|
||||
|
||||
assert not integration_session.in_transaction()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_separate_sessions_cas_topup_credit_exactly_once(
|
||||
integration_engine: AsyncEngine,
|
||||
) -> None:
|
||||
key_hash = uuid.uuid4().hex
|
||||
invoice = _lightning_invoice(
|
||||
purpose="topup",
|
||||
api_key_hash=key_hash,
|
||||
amount_sats=100,
|
||||
)
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=100_000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://mint:3338",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(key)
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
snapshot_a = _InvoiceSettlement.from_invoice(invoice)
|
||||
snapshot_b = _InvoiceSettlement.from_invoice(invoice)
|
||||
async with (
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as session_a,
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as session_b,
|
||||
):
|
||||
results = await asyncio.gather(
|
||||
_finalize_invoice_settlement(snapshot_a, session_a, 1_700_000_000),
|
||||
_finalize_invoice_settlement(snapshot_b, session_b, 1_700_000_001),
|
||||
)
|
||||
|
||||
assert sorted(settled for settled, _ in results) == [False, True]
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored_invoice = await verify.get(LightningInvoice, invoice.id)
|
||||
stored_key = await verify.get(ApiKey, key_hash)
|
||||
assert stored_invoice is not None
|
||||
assert stored_invoice.status == "paid"
|
||||
assert stored_key is not None
|
||||
assert stored_key.balance == 200_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_atomic_increment_preserves_concurrent_balance_mutation(
|
||||
integration_engine: AsyncEngine,
|
||||
) -> None:
|
||||
key_hash = uuid.uuid4().hex
|
||||
invoice = _lightning_invoice(
|
||||
purpose="topup", api_key_hash=key_hash, amount_sats=100
|
||||
)
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=100_000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://mint:3338",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(key)
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
async def debit_balance(session: AsyncSession) -> None:
|
||||
result = await session.exec( # type: ignore[call-overload]
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == key_hash)
|
||||
.values(balance=col(ApiKey.balance) - 10_000)
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
assert result.rowcount == 1
|
||||
await session.commit()
|
||||
|
||||
snapshot = _InvoiceSettlement.from_invoice(invoice)
|
||||
async with (
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as settlement,
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as debit,
|
||||
):
|
||||
settlement_result, _ = await asyncio.gather(
|
||||
_finalize_invoice_settlement(snapshot, settlement, 1_700_000_000),
|
||||
debit_balance(debit),
|
||||
)
|
||||
|
||||
assert settlement_result[0]
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored_key = await verify.get(ApiKey, key_hash)
|
||||
assert stored_key is not None
|
||||
assert stored_key.balance == 190_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_final_commit_rolls_back_claim_and_credit_for_retry(
|
||||
integration_engine: AsyncEngine,
|
||||
) -> None:
|
||||
key_hash = uuid.uuid4().hex
|
||||
invoice = _lightning_invoice(
|
||||
purpose="topup",
|
||||
api_key_hash=key_hash,
|
||||
amount_sats=100,
|
||||
)
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=100_000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://mint:3338",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(key)
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
snapshot = _InvoiceSettlement.from_invoice(invoice)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as failed:
|
||||
with patch.object(
|
||||
failed, "commit", AsyncMock(side_effect=Exception("db unavailable"))
|
||||
):
|
||||
with pytest.raises(Exception, match="db unavailable"):
|
||||
await _finalize_invoice_settlement(snapshot, failed, 1_700_000_000)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
pending = await verify.get(LightningInvoice, invoice.id)
|
||||
unchanged = await verify.get(ApiKey, key_hash)
|
||||
assert pending is not None
|
||||
assert pending.status == "pending"
|
||||
assert unchanged is not None
|
||||
assert unchanged.balance == 100_000
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as retry:
|
||||
settled, _ = await _finalize_invoice_settlement(
|
||||
snapshot, retry, 1_700_000_001
|
||||
)
|
||||
assert settled
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
paid = await verify.get(LightningInvoice, invoice.id)
|
||||
credited = await verify.get(ApiKey, key_hash)
|
||||
assert paid is not None
|
||||
assert paid.status == "paid"
|
||||
assert credited is not None
|
||||
assert credited.balance == 200_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_invoice_payment_retries_after_mint_success_and_db_failure(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
key_hash = uuid.uuid4().hex
|
||||
invoice = _lightning_invoice(
|
||||
purpose="topup", api_key_hash=key_hash, amount_sats=100
|
||||
)
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=100_000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://mint:3338",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(key)
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
wallet = Mock(
|
||||
proofs=[],
|
||||
keysets={"keyset-1": Mock()},
|
||||
load_proofs=AsyncMock(),
|
||||
get_mint_quote=AsyncMock(return_value=Mock(paid=True)),
|
||||
restore_tokens_for_keyset=AsyncMock(),
|
||||
)
|
||||
|
||||
async def mint(amount: int, quote_id: str) -> list[Proof]:
|
||||
proofs = [Proof(amount=amount, mint_id=quote_id)]
|
||||
wallet.proofs.extend(proofs)
|
||||
return proofs
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=mint)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as failed:
|
||||
stored = await failed.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch(
|
||||
"routstr.lightning._finalize_invoice_settlement",
|
||||
AsyncMock(side_effect=Exception("db unavailable")),
|
||||
),
|
||||
):
|
||||
await check_invoice_payment(stored, failed)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
pending = await verify.get(LightningInvoice, invoice.id)
|
||||
unchanged = await verify.get(ApiKey, key_hash)
|
||||
assert pending is not None
|
||||
assert pending.status == "settlement_pending"
|
||||
assert unchanged is not None
|
||||
assert unchanged.balance == 100_000
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as retry:
|
||||
stored = await retry.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
await check_invoice_payment(stored, retry)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
paid = await verify.get(LightningInvoice, invoice.id)
|
||||
credited = await verify.get(ApiKey, key_hash)
|
||||
assert paid is not None
|
||||
assert paid.status == "paid"
|
||||
assert credited is not None
|
||||
assert credited.balance == 200_000
|
||||
|
||||
wallet.mint.assert_awaited_once_with(100, quote_id=invoice.payment_hash)
|
||||
wallet.restore_tokens_for_keyset.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_expiry_cas_cannot_overwrite_concurrent_paid_invoice(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _lightning_invoice(expires_at=0)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as caller:
|
||||
stale = await caller.get(LightningInvoice, invoice.id)
|
||||
assert stale is not None
|
||||
await caller.commit()
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as paid:
|
||||
result = await paid.exec( # type: ignore[call-overload]
|
||||
update(LightningInvoice)
|
||||
.where(col(LightningInvoice.id) == invoice.id)
|
||||
.values(status="paid", paid_at=123)
|
||||
)
|
||||
assert result.rowcount == 1
|
||||
await paid.commit()
|
||||
|
||||
expired = await _expire_invoice_if_authoritatively_unpaid(
|
||||
stale, caller, True
|
||||
)
|
||||
|
||||
assert expired is False
|
||||
assert stale.status == "paid"
|
||||
assert stale.paid_at == 123
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
assert stored.status == "paid"
|
||||
assert stored.paid_at == 123
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_paid_quote_worker_does_not_mint_after_expiry_claim_wins(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _lightning_invoice(expires_at=0)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
quote_started = asyncio.Event()
|
||||
release_quote = asyncio.Event()
|
||||
|
||||
async def paid_quote_after_expiry(*_args: object, **_kwargs: object) -> Mock:
|
||||
quote_started.set()
|
||||
await release_quote.wait()
|
||||
return Mock(paid=True)
|
||||
|
||||
wallet = Mock(
|
||||
get_mint_quote=AsyncMock(side_effect=paid_quote_after_expiry),
|
||||
mint=AsyncMock(),
|
||||
)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as worker:
|
||||
observed_pending = await worker.get(LightningInvoice, invoice.id)
|
||||
assert observed_pending is not None
|
||||
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
settlement_task = asyncio.create_task(
|
||||
check_invoice_payment(observed_pending, worker)
|
||||
)
|
||||
await quote_started.wait()
|
||||
|
||||
async with AsyncSession(
|
||||
integration_engine, expire_on_commit=False
|
||||
) as expirer:
|
||||
expiry_view = await expirer.get(LightningInvoice, invoice.id)
|
||||
assert expiry_view is not None
|
||||
await expirer.commit()
|
||||
assert await _expire_invoice_if_authoritatively_unpaid(
|
||||
expiry_view, expirer, True
|
||||
)
|
||||
|
||||
release_quote.set()
|
||||
assert await settlement_task is False
|
||||
|
||||
wallet.mint.assert_not_awaited()
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
assert stored.status == "expired"
|
||||
@@ -0,0 +1,178 @@
|
||||
"""Money-safety regression coverage for automatic wallet payouts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable, Coroutine
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core import db
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.core.settings import settings
|
||||
from routstr.wallet import credit_balance, periodic_payout
|
||||
|
||||
PRIMARY_MINT = "http://primary:3338"
|
||||
REFUND_MINT = "http://refund:3338"
|
||||
PAYOUT_INTERVAL = 987
|
||||
|
||||
|
||||
class _LoopBreak(Exception):
|
||||
"""Stop the otherwise-infinite payout loop after one cycle."""
|
||||
|
||||
|
||||
def _one_payout_cycle() -> Callable[[float], Coroutine[Any, Any, None]]:
|
||||
intervals_seen = 0
|
||||
|
||||
async def sleep(seconds: float) -> None:
|
||||
nonlocal intervals_seen
|
||||
if seconds == PAYOUT_INTERVAL:
|
||||
intervals_seen += 1
|
||||
if intervals_seen == 2:
|
||||
raise _LoopBreak()
|
||||
|
||||
return sleep
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cross_mint_liability_is_not_paid_as_owner_profit(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Refund preferences must not make primary-mint customer funds payable."""
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(
|
||||
ApiKey(
|
||||
hashed_key="cross-mint-key",
|
||||
balance=50_000,
|
||||
refund_mint_url=REFUND_MINT,
|
||||
refund_currency="sat",
|
||||
)
|
||||
)
|
||||
await setup.commit()
|
||||
|
||||
primary_proof = MagicMock(amount=50)
|
||||
raw_send = AsyncMock(return_value=50)
|
||||
|
||||
def proofs_for_mint(
|
||||
_wallet: object, mint_url: str, unit: str, **_kwargs: object
|
||||
) -> list[MagicMock]:
|
||||
if mint_url == PRIMARY_MINT and unit == "sat":
|
||||
return [primary_proof]
|
||||
return []
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", [REFUND_MINT]),
|
||||
patch.object(settings, "primary_mint", PRIMARY_MINT),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.test"),
|
||||
patch.object(settings, "payout_interval_seconds", PAYOUT_INTERVAL),
|
||||
patch.object(settings, "min_payout_sat", 10),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_payout_cycle()),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(side_effect=proofs_for_mint),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, _wallet: proofs),
|
||||
),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
|
||||
):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
raw_send.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_payout_does_not_send_proofs_whose_liability_commit_is_in_flight(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Proof visibility before liability commit must not expose customer funds."""
|
||||
key = ApiKey(
|
||||
hashed_key="in-flight-topup-key",
|
||||
balance=0,
|
||||
refund_mint_url=PRIMARY_MINT,
|
||||
refund_currency="sat",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(key)
|
||||
await setup.commit()
|
||||
|
||||
proofs: list[MagicMock] = []
|
||||
proof_visible = asyncio.Event()
|
||||
finish_redemption = asyncio.Event()
|
||||
liability_read = asyncio.Event()
|
||||
|
||||
async def redeem_token(
|
||||
token: str,
|
||||
destination_mint: str | None = None,
|
||||
destination_unit: str | None = None,
|
||||
) -> tuple[int, str, str]:
|
||||
proofs.append(MagicMock(amount=200))
|
||||
proof_visible.set()
|
||||
await finish_redemption.wait()
|
||||
return 200, "sat", PRIMARY_MINT
|
||||
|
||||
real_total_liability = db.total_user_liability
|
||||
|
||||
async def read_liability(_session: AsyncSession) -> int:
|
||||
async with db.create_session() as snapshot_session:
|
||||
value = await real_total_liability(snapshot_session)
|
||||
liability_read.set()
|
||||
return value
|
||||
|
||||
raw_send = AsyncMock(return_value=200)
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", PRIMARY_MINT),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.test"),
|
||||
patch.object(settings, "payout_interval_seconds", PAYOUT_INTERVAL),
|
||||
patch.object(settings, "min_payout_sat", 10),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_payout_cycle()),
|
||||
patch("routstr.wallet.recieve_token", AsyncMock(side_effect=redeem_token)),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(side_effect=lambda *_args, **_kwargs: list(proofs)),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda visible, _wallet: visible),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.total_user_liability",
|
||||
AsyncMock(side_effect=read_liability),
|
||||
),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
|
||||
):
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as credit_session:
|
||||
stored_key = await credit_session.get(ApiKey, key.hashed_key)
|
||||
assert stored_key is not None
|
||||
credit_task = asyncio.create_task(
|
||||
credit_balance("cashu-token", stored_key, credit_session)
|
||||
)
|
||||
await asyncio.wait_for(proof_visible.wait(), timeout=2)
|
||||
|
||||
payout_task = asyncio.create_task(periodic_payout())
|
||||
try:
|
||||
await asyncio.wait_for(liability_read.wait(), timeout=0.1)
|
||||
liability_was_read_while_crediting = True
|
||||
except TimeoutError:
|
||||
liability_was_read_while_crediting = False
|
||||
|
||||
finish_redemption.set()
|
||||
await asyncio.wait_for(credit_task, timeout=2)
|
||||
|
||||
with pytest.raises(_LoopBreak):
|
||||
await asyncio.wait_for(payout_task, timeout=2)
|
||||
|
||||
assert liability_was_read_while_crediting is False
|
||||
raw_send.assert_not_awaited()
|
||||
@@ -0,0 +1,602 @@
|
||||
"""Real-database tests for the PPQ auto top-up claim lifecycle.
|
||||
|
||||
These exercise the claim against actual SQL rather than mocked sessions,
|
||||
because the guarantees under test are all about what the database will and
|
||||
will not let two concurrent writers do.
|
||||
"""
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlmodel import select
|
||||
|
||||
from routstr.core.db import CashuTransaction, create_session
|
||||
from routstr.upstream.auto_topup import (
|
||||
PPQ_PHASE_CLAIMED,
|
||||
PPQ_PHASE_IN_FLIGHT,
|
||||
PPQ_PHASE_RECONCILE,
|
||||
_claim_ppq_topup,
|
||||
_ppq_payment_id,
|
||||
_ppq_payment_usd,
|
||||
_ppq_request_id,
|
||||
_ppq_spent_last_24h_usd,
|
||||
_ppq_state_id_for_provider,
|
||||
_record_ppq_invoice,
|
||||
_set_ppq_state_terminal,
|
||||
get_ppq_auto_topup_state,
|
||||
release_ppq_auto_topup_state,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
def _row(provider_id: int = 1) -> MagicMock:
|
||||
row = MagicMock()
|
||||
row.id = provider_id
|
||||
return row
|
||||
|
||||
|
||||
async def _seed_provider(provider_id: int = 1, slug: str = "ppq") -> None:
|
||||
"""Claim creation is fenced on the provider row existing; seed it."""
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
|
||||
async with create_session() as session:
|
||||
session.add(
|
||||
UpstreamProviderRow(
|
||||
id=provider_id,
|
||||
slug=slug,
|
||||
provider_type="ppqai",
|
||||
base_url="https://api.ppq.ai",
|
||||
api_key="secret",
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def _state_row(provider_id: int = 1) -> CashuTransaction | None:
|
||||
async with create_session() as session:
|
||||
return await session.get(
|
||||
CashuTransaction, _ppq_state_id_for_provider(provider_id)
|
||||
)
|
||||
|
||||
|
||||
async def _seed_claim(
|
||||
provider_id: int,
|
||||
phase: str,
|
||||
invoice_id: str,
|
||||
lease_expires_at: int,
|
||||
quote_id: str = "quote-1",
|
||||
) -> str:
|
||||
"""Seed a claim row and return its state token (the full request_id)."""
|
||||
token = _ppq_request_id(
|
||||
"operation-1", lease_expires_at, phase, invoice_id, quote_id
|
||||
)
|
||||
async with create_session() as session:
|
||||
session.add(
|
||||
CashuTransaction(
|
||||
id=_ppq_state_id_for_provider(provider_id),
|
||||
token="lnbc-invoice",
|
||||
amount=102,
|
||||
unit="sat",
|
||||
type="out",
|
||||
request_id=token,
|
||||
mint_url="https://mint.test",
|
||||
collected=False,
|
||||
source="ppq_auto_topup",
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
return token
|
||||
|
||||
|
||||
async def test_second_claim_is_refused_while_the_first_is_active(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
await _seed_provider()
|
||||
assert await _claim_ppq_topup(_row()) is not None
|
||||
# The whole point of the claim: a concurrent cycle must not get one.
|
||||
assert await _claim_ppq_topup(_row()) is None
|
||||
|
||||
async with create_session() as session:
|
||||
rows = (await session.exec(select(CashuTransaction))).all()
|
||||
assert len(rows) == 1
|
||||
|
||||
|
||||
async def test_claim_is_reusable_once_the_previous_attempt_finished(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
await _seed_provider()
|
||||
first = await _claim_ppq_topup(_row())
|
||||
assert first is not None
|
||||
assert await _set_ppq_state_terminal(_row(), first, collected=True, swept=False)
|
||||
|
||||
second = await _claim_ppq_topup(_row())
|
||||
assert second is not None and second != first
|
||||
|
||||
|
||||
async def test_recording_the_invoice_moves_the_claim_in_flight(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
await _seed_provider()
|
||||
operation_id = await _claim_ppq_topup(_row())
|
||||
assert operation_id is not None
|
||||
|
||||
state = await get_ppq_auto_topup_state(1)
|
||||
assert state["phase"] == PPQ_PHASE_CLAIMED
|
||||
assert state["releasable"] is True
|
||||
assert state["invoice_id"] is None
|
||||
|
||||
lease = await _record_ppq_invoice(
|
||||
_row(),
|
||||
operation_id,
|
||||
invoice="lnbc-invoice",
|
||||
invoice_id="invoice-1",
|
||||
quote_id="quote-1",
|
||||
amount=102,
|
||||
amount_usd=10,
|
||||
unit="sat",
|
||||
mint_url="https://mint.test",
|
||||
)
|
||||
assert lease > int(time.time())
|
||||
|
||||
state = await get_ppq_auto_topup_state(1)
|
||||
assert state["phase"] == PPQ_PHASE_IN_FLIGHT
|
||||
assert state["invoice_id"] == "invoice-1"
|
||||
# A payment is committed to a mint, so an admin must not sweep it.
|
||||
assert state["releasable"] is False
|
||||
# The raw BOLT11 invoice must never reach the admin API.
|
||||
assert "token" not in state
|
||||
|
||||
|
||||
async def test_release_refuses_an_in_flight_claim(patched_db_engine: Any) -> None:
|
||||
token = await _seed_claim(
|
||||
1, PPQ_PHASE_IN_FLIGHT, "invoice-1", int(time.time()) + 900
|
||||
)
|
||||
|
||||
outcome = await release_ppq_auto_topup_state(1, state_token=token)
|
||||
|
||||
assert outcome.released is False
|
||||
assert outcome.reason == "payment_in_flight"
|
||||
row = await _state_row()
|
||||
assert row is not None and row.swept is False
|
||||
|
||||
|
||||
async def test_release_refuses_a_stale_state_token(patched_db_engine: Any) -> None:
|
||||
await _seed_claim(1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) + 900)
|
||||
|
||||
outcome = await release_ppq_auto_topup_state(1, state_token="ppq:stale:token")
|
||||
|
||||
assert outcome.released is False
|
||||
assert outcome.reason == "stale_state"
|
||||
row = await _state_row()
|
||||
assert row is not None and row.swept is False
|
||||
|
||||
|
||||
async def test_release_accepts_a_reconcile_claim(patched_db_engine: Any) -> None:
|
||||
token = await _seed_claim(
|
||||
1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) + 900
|
||||
)
|
||||
|
||||
outcome = await release_ppq_auto_topup_state(1, state_token=token)
|
||||
|
||||
assert outcome.released is True
|
||||
row = await _state_row()
|
||||
assert row is not None and row.swept is True
|
||||
|
||||
|
||||
async def test_expired_in_flight_claim_becomes_releasable(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
# A worker that died mid-payment must not lock the provider forever.
|
||||
token = await _seed_claim(1, PPQ_PHASE_IN_FLIGHT, "invoice-1", int(time.time()) - 1)
|
||||
|
||||
assert (await get_ppq_auto_topup_state(1))["releasable"] is True
|
||||
outcome = await release_ppq_auto_topup_state(1, state_token=token)
|
||||
assert outcome.released is True
|
||||
|
||||
|
||||
async def test_release_reports_no_active_claim_once_swept(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
token = await _seed_claim(
|
||||
1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) + 900
|
||||
)
|
||||
assert (await release_ppq_auto_topup_state(1, state_token=token)).released
|
||||
|
||||
outcome = await release_ppq_auto_topup_state(1, state_token=token)
|
||||
assert outcome.released is False
|
||||
assert outcome.reason == "no_active_claim"
|
||||
|
||||
|
||||
async def test_terminal_write_fails_after_the_claim_was_released(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
"""The symptom an admin release leaves behind for the owning worker."""
|
||||
token = await _seed_claim(
|
||||
1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) + 900
|
||||
)
|
||||
assert (await release_ppq_auto_topup_state(1, state_token=token)).released
|
||||
|
||||
assert (
|
||||
await _set_ppq_state_terminal(
|
||||
_row(), "operation-1", collected=True, swept=False
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
async def test_ppq_claim_rows_are_excluded_from_the_admin_transaction_list(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
from routstr.core.admin import get_transactions_api
|
||||
|
||||
await _seed_provider()
|
||||
await _claim_ppq_topup(_row())
|
||||
async with create_session() as session:
|
||||
session.add(
|
||||
CashuTransaction(
|
||||
id="real-transaction",
|
||||
token="cashuAreal",
|
||||
amount=50,
|
||||
unit="sat",
|
||||
type="out",
|
||||
source="x-cashu",
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
result = await get_transactions_api()
|
||||
|
||||
ids = {t["id"] for t in result["transactions"]} # type: ignore[index,union-attr]
|
||||
assert "real-transaction" in ids
|
||||
assert _ppq_state_id_for_provider(1) not in ids
|
||||
|
||||
|
||||
async def test_ppq_payment_audit_row_is_visible_and_survives_next_claim(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
from routstr.core.admin import get_transactions_api
|
||||
|
||||
await _seed_provider()
|
||||
operation_id = await _claim_ppq_topup(_row())
|
||||
assert operation_id is not None
|
||||
await _record_ppq_invoice(
|
||||
_row(),
|
||||
operation_id,
|
||||
invoice="lnbc-secret-invoice",
|
||||
invoice_id="invoice-1",
|
||||
quote_id="quote-1",
|
||||
amount=102,
|
||||
amount_usd=10,
|
||||
unit="sat",
|
||||
mint_url="https://mint.test",
|
||||
)
|
||||
assert await _set_ppq_state_terminal(
|
||||
_row(), operation_id, collected=True, swept=False
|
||||
)
|
||||
|
||||
result = await get_transactions_api(source="ppq_auto_topup")
|
||||
transactions = result["transactions"]
|
||||
assert len(transactions) == 1
|
||||
audit = transactions[0]
|
||||
assert audit["id"] == _ppq_payment_id(operation_id)
|
||||
assert audit["token"] == "ppq-invoice:invoice-1:usd:10"
|
||||
assert audit["collected"] is True
|
||||
assert "lnbc-secret-invoice" not in audit["token"]
|
||||
|
||||
# Reusing the deterministic claim lock must not overwrite history.
|
||||
assert await _claim_ppq_topup(_row()) is not None
|
||||
async with create_session() as session:
|
||||
assert await session.get(CashuTransaction, audit["id"]) is not None
|
||||
|
||||
|
||||
async def test_reconcile_settles_a_recorded_invoice(patched_db_engine: Any) -> None:
|
||||
from routstr.upstream.auto_topup import _reconcile_ppq_state
|
||||
|
||||
await _seed_claim(1, PPQ_PHASE_IN_FLIGHT, "invoice-1", int(time.time()) + 900)
|
||||
provider = MagicMock()
|
||||
provider.check_topup_status = AsyncMock(return_value=True)
|
||||
|
||||
# Still suppresses this cycle, but the claim is now finished.
|
||||
assert await _reconcile_ppq_state(_row(), provider) is True
|
||||
|
||||
row = await _state_row()
|
||||
assert row is not None and row.collected is True
|
||||
|
||||
|
||||
async def test_stale_token_from_before_a_phase_change_cannot_release(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
"""The blocker scenario: admin reviews `claimed`, payment turns ambiguous.
|
||||
|
||||
The operation id is identical in both states, so an id-based fence would
|
||||
let the stale confirmation land. The full state token must not.
|
||||
"""
|
||||
await _seed_provider()
|
||||
operation_id = await _claim_ppq_topup(_row())
|
||||
assert operation_id is not None
|
||||
reviewed = await get_ppq_auto_topup_state(1)
|
||||
assert reviewed["phase"] == PPQ_PHASE_CLAIMED
|
||||
|
||||
# Worker records the invoice: same operation, new phase, proofs committed.
|
||||
await _record_ppq_invoice(
|
||||
_row(),
|
||||
operation_id,
|
||||
invoice="lnbc-invoice",
|
||||
invoice_id="invoice-1",
|
||||
quote_id="quote-1",
|
||||
amount=102,
|
||||
amount_usd=10,
|
||||
unit="sat",
|
||||
mint_url="https://mint.test",
|
||||
)
|
||||
|
||||
outcome = await release_ppq_auto_topup_state(
|
||||
1, state_token=str(reviewed["state_token"])
|
||||
)
|
||||
assert outcome.released is False
|
||||
assert outcome.reason == "stale_state"
|
||||
row = await _state_row()
|
||||
assert row is not None and row.swept is False
|
||||
|
||||
|
||||
async def test_concurrent_claims_only_one_wins(patched_db_engine: Any) -> None:
|
||||
import asyncio
|
||||
|
||||
await _seed_provider()
|
||||
|
||||
results = await asyncio.gather(
|
||||
*(_claim_ppq_topup(_row()) for _ in range(5)), return_exceptions=True
|
||||
)
|
||||
winners = [r for r in results if isinstance(r, str)]
|
||||
assert len(winners) == 1
|
||||
|
||||
async with create_session() as session:
|
||||
rows = (await session.exec(select(CashuTransaction))).all()
|
||||
assert len(rows) == 1
|
||||
|
||||
|
||||
async def test_reconcile_releases_claim_when_mint_reports_unpaid(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
from routstr.upstream.auto_topup import _reconcile_ppq_state
|
||||
|
||||
# Lease expired, PPQ never credited: only the mint's own "unpaid" answer
|
||||
# may hand the claim back.
|
||||
await _seed_claim(1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) - 1)
|
||||
provider = MagicMock()
|
||||
provider.check_topup_status = AsyncMock(return_value=False)
|
||||
|
||||
with patch(
|
||||
"routstr.upstream.auto_topup.check_bolt11_payment_status",
|
||||
AsyncMock(return_value="unpaid"),
|
||||
) as status:
|
||||
suppressed = await _reconcile_ppq_state(_row(), provider)
|
||||
|
||||
status.assert_awaited_once_with("https://mint.test", "sat", "quote-1")
|
||||
assert suppressed is False
|
||||
row = await _state_row()
|
||||
assert row is not None and row.swept is True
|
||||
|
||||
|
||||
async def test_reconcile_keeps_claim_when_mint_answer_is_not_final(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
from routstr.upstream.auto_topup import _reconcile_ppq_state
|
||||
|
||||
await _seed_claim(1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) - 1)
|
||||
provider = MagicMock()
|
||||
provider.check_topup_status = AsyncMock(return_value=False)
|
||||
|
||||
for answer in ("paid", "pending", "unknown"):
|
||||
with patch(
|
||||
"routstr.upstream.auto_topup.check_bolt11_payment_status",
|
||||
AsyncMock(return_value=answer),
|
||||
):
|
||||
assert await _reconcile_ppq_state(_row(), provider) is True
|
||||
row = await _state_row()
|
||||
assert row is not None and row.swept is False, answer
|
||||
|
||||
|
||||
async def test_release_endpoint_maps_refusals_to_409(patched_db_engine: Any) -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.core.admin import (
|
||||
ReleasePPQAutoTopupRequest,
|
||||
release_ppq_auto_topup_api,
|
||||
)
|
||||
|
||||
provider_row = MagicMock()
|
||||
provider_row.provider_type = "ppqai"
|
||||
|
||||
token = await _seed_claim(
|
||||
1, PPQ_PHASE_IN_FLIGHT, "invoice-1", int(time.time()) + 900
|
||||
)
|
||||
|
||||
with patch(
|
||||
"routstr.core.admin._require_ppq_provider",
|
||||
AsyncMock(return_value=provider_row),
|
||||
):
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await release_ppq_auto_topup_api(
|
||||
1,
|
||||
ReleasePPQAutoTopupRequest(
|
||||
confirmed_safe_to_retry=True, state_token=token
|
||||
),
|
||||
)
|
||||
assert excinfo.value.status_code == 409
|
||||
assert "in flight" in excinfo.value.detail
|
||||
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await release_ppq_auto_topup_api(
|
||||
1,
|
||||
ReleasePPQAutoTopupRequest(
|
||||
confirmed_safe_to_retry=True, state_token="ppq:wrong"
|
||||
),
|
||||
)
|
||||
assert excinfo.value.status_code == 409
|
||||
assert "changed since" in excinfo.value.detail
|
||||
|
||||
|
||||
async def test_provider_delete_is_blocked_by_an_active_claim(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.core.admin import delete_upstream_provider
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
|
||||
async with create_session() as session:
|
||||
session.add(
|
||||
UpstreamProviderRow(
|
||||
id=1,
|
||||
slug="ppq",
|
||||
provider_type="ppqai",
|
||||
base_url="https://api.ppq.ai",
|
||||
api_key="secret",
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
await _seed_claim(1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) + 900)
|
||||
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await delete_upstream_provider("1")
|
||||
assert excinfo.value.status_code == 409
|
||||
|
||||
# Provider must still exist.
|
||||
async with create_session() as session:
|
||||
assert await session.get(UpstreamProviderRow, 1) is not None
|
||||
|
||||
|
||||
async def test_claim_is_refused_when_the_provider_row_is_gone(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
"""The worker's half of the delete race: no provider row, no claim."""
|
||||
assert await _claim_ppq_topup(_row()) is None
|
||||
|
||||
async with create_session() as session:
|
||||
rows = (await session.exec(select(CashuTransaction))).all()
|
||||
assert rows == []
|
||||
|
||||
|
||||
async def test_claim_is_refused_after_a_provider_type_change(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
|
||||
await _seed_provider()
|
||||
async with create_session() as session:
|
||||
provider = await session.get(UpstreamProviderRow, 1)
|
||||
assert provider is not None
|
||||
provider.provider_type = "openai"
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
|
||||
assert await _claim_ppq_topup(_row()) is None
|
||||
|
||||
|
||||
async def test_disabled_provider_with_claim_still_reconciles(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
"""A claim tracks committed money; eligibility must not stop reconciling."""
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
from routstr.upstream.auto_topup import _reconcile_all_ppq_claims
|
||||
|
||||
await _seed_provider()
|
||||
async with create_session() as session:
|
||||
provider = await session.get(UpstreamProviderRow, 1)
|
||||
assert provider is not None
|
||||
provider.enabled = False
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
await _seed_claim(1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) + 900)
|
||||
|
||||
ppq = MagicMock()
|
||||
ppq.check_topup_status = AsyncMock(return_value=True)
|
||||
with patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=ppq,
|
||||
):
|
||||
await _reconcile_all_ppq_claims()
|
||||
|
||||
row = await _state_row()
|
||||
assert row is not None and row.collected is True
|
||||
|
||||
|
||||
async def test_claim_without_api_key_still_reconciles_via_the_mint(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
from routstr.upstream.auto_topup import _reconcile_all_ppq_claims
|
||||
|
||||
await _seed_provider()
|
||||
async with create_session() as session:
|
||||
provider = await session.get(UpstreamProviderRow, 1)
|
||||
assert provider is not None
|
||||
provider.api_key = ""
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
# Lease expired, so the mint may be consulted.
|
||||
await _seed_claim(1, PPQ_PHASE_RECONCILE, "invoice-1", int(time.time()) - 1)
|
||||
|
||||
with patch(
|
||||
"routstr.upstream.auto_topup.check_bolt11_payment_status",
|
||||
AsyncMock(return_value="unpaid"),
|
||||
) as status:
|
||||
await _reconcile_all_ppq_claims()
|
||||
|
||||
# No API key: PPQ was never polled, but the mint was, and its definitive
|
||||
# "unpaid" released the claim.
|
||||
status.assert_awaited_once()
|
||||
row = await _state_row()
|
||||
assert row is not None and row.swept is True
|
||||
|
||||
|
||||
def test_ppq_payment_usd_prefers_stamped_amount() -> None:
|
||||
# Stamped rows must not move with the BTC price.
|
||||
assert _ppq_payment_usd(102, "sat", "ppq-invoice:a:usd:10", 0.5) == 10.0
|
||||
|
||||
|
||||
def test_ppq_payment_usd_falls_back_to_current_price() -> None:
|
||||
# Rows recorded before the stamp existed convert sats at today's price.
|
||||
assert _ppq_payment_usd(2000, "sat", "ppq-invoice:legacy", 0.001) == 2.0
|
||||
assert _ppq_payment_usd(2_000_000, "msat", "ppq-invoice:legacy", 0.001) == 2.0
|
||||
|
||||
|
||||
def test_ppq_payment_usd_survives_malformed_stamp() -> None:
|
||||
assert _ppq_payment_usd(3000, "sat", "ppq-invoice:x:usd:oops", 0.001) == 3.0
|
||||
|
||||
|
||||
async def test_daily_spend_ignores_provably_unattempted_payments(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
def _payment(
|
||||
id_: str, token: str, collected: bool, swept: bool
|
||||
) -> CashuTransaction:
|
||||
return CashuTransaction(
|
||||
id=id_,
|
||||
token=token,
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
source="ppq_auto_topup",
|
||||
collected=collected,
|
||||
swept=swept,
|
||||
)
|
||||
|
||||
async with create_session() as session:
|
||||
# Settled, in-flight, and provably-unattempted payments plus a
|
||||
# pre-stamp row: only the unattempted one must be excluded.
|
||||
session.add(_payment("pay-usd-1", "ppq-invoice:a:usd:100", True, False))
|
||||
session.add(_payment("pay-usd-2", "ppq-invoice:b:usd:50", False, False))
|
||||
session.add(_payment("pay-usd-3", "ppq-invoice:c:usd:25", False, True))
|
||||
legacy = _payment("pay-usd-4", "ppq-invoice:legacy", True, False)
|
||||
legacy.amount = 2000
|
||||
session.add(legacy)
|
||||
await session.commit()
|
||||
|
||||
assert await _ppq_spent_last_24h_usd(0.001) == 152.0
|
||||
@@ -0,0 +1,65 @@
|
||||
"""Integration coverage for proxy database-session lifetime."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import Response
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticated_proxy_releases_db_connection_before_upstream_headers(
|
||||
integration_engine: AsyncEngine,
|
||||
integration_session: AsyncSession,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Slow upstream header waits must not retain a checked-out DB connection."""
|
||||
key = ApiKey(
|
||||
hashed_key="proxy-pool-key",
|
||||
balance=1_000_000,
|
||||
refund_mint_url="http://primary:3338",
|
||||
refund_currency="sat",
|
||||
)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = {"authorization": "Bearer test-key"}
|
||||
request.body = AsyncMock(return_value=json.dumps({"model": "test-model"}).encode())
|
||||
request.url.path = "/v1/chat/completions"
|
||||
request.state.request_id = "pool-hold-regression"
|
||||
|
||||
model = MagicMock()
|
||||
upstream = MagicMock()
|
||||
upstream.provider_type = "test"
|
||||
upstream.prepare_headers.return_value = {}
|
||||
|
||||
async def wait_for_headers(*args: object, **kwargs: object) -> Response:
|
||||
assert integration_engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
return Response(status_code=200)
|
||||
|
||||
upstream.forward_request = AsyncMock(side_effect=wait_for_headers)
|
||||
|
||||
with (
|
||||
patch("routstr.proxy.get_candidates", return_value=[(model, upstream)]),
|
||||
patch("routstr.proxy.get_max_cost_for_model", AsyncMock(return_value=100)),
|
||||
patch(
|
||||
"routstr.proxy.calculate_discounted_max_cost",
|
||||
AsyncMock(return_value=100),
|
||||
),
|
||||
patch("routstr.proxy.check_token_balance"),
|
||||
patch("routstr.proxy.get_bearer_token_key", AsyncMock(return_value=key)),
|
||||
):
|
||||
response = await proxy_module._proxy(
|
||||
request, "v1/chat/completions", integration_session
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
@@ -126,8 +126,11 @@ async def test_parent_and_child_keys_are_not_pruned(
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pending_invoice_protects_key(patched_db_engine: None) -> None:
|
||||
"""A key referenced by a pending topup invoice is never pruned mid-topup."""
|
||||
@pytest.mark.parametrize("status", ["pending", "settlement_pending"])
|
||||
async def test_retryable_invoice_protects_key(
|
||||
patched_db_engine: None, status: str
|
||||
) -> None:
|
||||
"""A key referenced by a retryable topup invoice is never pruned mid-topup."""
|
||||
key = _dead_key(LONG_AGO)
|
||||
invoice = LightningInvoice(
|
||||
id=f"inv_{uuid.uuid4().hex}",
|
||||
@@ -135,7 +138,7 @@ async def test_pending_invoice_protects_key(patched_db_engine: None) -> None:
|
||||
amount_sats=10,
|
||||
description="topup",
|
||||
payment_hash=uuid.uuid4().hex,
|
||||
status="pending",
|
||||
status=status,
|
||||
api_key_hash=key.hashed_key,
|
||||
purpose="topup",
|
||||
expires_at=NOW + 10_000,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -141,7 +141,7 @@ async def test_revert_with_zero_reserved_balance_is_noop(
|
||||
Previously this would drive reserved_balance negative. With the floor guard,
|
||||
it should return False and leave reserved_balance at 0.
|
||||
"""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
from routstr.auth import pay_for_request, revert_pay_for_request
|
||||
|
||||
unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
@@ -151,8 +151,12 @@ async def test_revert_with_zero_reserved_balance_is_noop(
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
await pay_for_request(test_key, 100, integration_session)
|
||||
test_key.reserved_balance = 0
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Try to revert more than available — should be a no-op
|
||||
# A stale cleanup already released the aggregate reservation.
|
||||
result = await revert_pay_for_request(test_key, integration_session, 100)
|
||||
|
||||
await integration_session.refresh(test_key)
|
||||
@@ -161,8 +165,8 @@ async def test_revert_with_zero_reserved_balance_is_noop(
|
||||
assert test_key.reserved_balance == 0, (
|
||||
f"Reserved balance should remain 0, got: {test_key.reserved_balance}"
|
||||
)
|
||||
assert test_key.total_requests == 0, (
|
||||
f"Total requests should remain 0, got: {test_key.total_requests}"
|
||||
assert test_key.total_requests == 1, (
|
||||
f"Total requests should remain 1, got: {test_key.total_requests}"
|
||||
)
|
||||
|
||||
|
||||
@@ -171,17 +175,18 @@ async def test_revert_with_sufficient_reserved_balance_succeeds(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test that revert_pay_for_request works correctly when there is enough reserved balance."""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
from routstr.auth import pay_for_request, revert_pay_for_request
|
||||
|
||||
unique_key = f"test_revert_ok_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
balance=5000,
|
||||
reserved_balance=500,
|
||||
total_requests=3,
|
||||
reserved_balance=0,
|
||||
total_requests=2,
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
await pay_for_request(test_key, 500, integration_session)
|
||||
|
||||
result = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
|
||||
@@ -202,17 +207,21 @@ async def test_revert_partial_reserved_balance_is_noop(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test that reverting more than the current reserved_balance is a no-op."""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
from routstr.auth import pay_for_request, revert_pay_for_request
|
||||
|
||||
unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
balance=5000,
|
||||
reserved_balance=50,
|
||||
total_requests=1,
|
||||
reserved_balance=0,
|
||||
total_requests=0,
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
await pay_for_request(test_key, 500, integration_session)
|
||||
test_key.reserved_balance = 50
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Try to revert 500 when only 50 is reserved — should be no-op
|
||||
result = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
@@ -237,20 +246,28 @@ async def test_double_revert_prevented(
|
||||
This simulates the double-revert scenario where both upstream/base.py
|
||||
and proxy.py attempt to revert the same reservation.
|
||||
"""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
from routstr.auth import (
|
||||
get_reservation_snapshot,
|
||||
pay_for_request,
|
||||
revert_pay_for_request,
|
||||
)
|
||||
|
||||
unique_key = f"test_double_revert_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
balance=10000,
|
||||
reserved_balance=500,
|
||||
total_requests=5,
|
||||
reserved_balance=0,
|
||||
total_requests=4,
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
await pay_for_request(test_key, 500, integration_session)
|
||||
snapshot = await get_reservation_snapshot(test_key, integration_session)
|
||||
|
||||
# First revert — should succeed
|
||||
result1 = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
result1 = await revert_pay_for_request(
|
||||
test_key, integration_session, 500, snapshot
|
||||
)
|
||||
await integration_session.refresh(test_key)
|
||||
|
||||
assert result1 is True
|
||||
@@ -258,7 +275,9 @@ async def test_double_revert_prevented(
|
||||
assert test_key.total_requests == 4
|
||||
|
||||
# Second revert of the same amount — should be no-op
|
||||
result2 = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
result2 = await revert_pay_for_request(
|
||||
test_key, integration_session, 500, snapshot
|
||||
)
|
||||
await integration_session.refresh(test_key)
|
||||
|
||||
assert result2 is False, "Second revert should be a no-op"
|
||||
@@ -279,22 +298,30 @@ async def test_sequential_reverts_never_go_negative(
|
||||
Simulates the double-revert scenario where multiple code paths
|
||||
attempt to revert the same reservation.
|
||||
"""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
from routstr.auth import (
|
||||
get_reservation_snapshot,
|
||||
pay_for_request,
|
||||
revert_pay_for_request,
|
||||
)
|
||||
|
||||
unique_key = f"test_multi_revert_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
balance=10000,
|
||||
reserved_balance=500,
|
||||
total_requests=5,
|
||||
reserved_balance=0,
|
||||
total_requests=4,
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
await pay_for_request(test_key, 500, integration_session)
|
||||
snapshot = await get_reservation_snapshot(test_key, integration_session)
|
||||
|
||||
# Run 5 sequential reverts for the same 500 reservation
|
||||
results = []
|
||||
for _ in range(5):
|
||||
r = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
r = await revert_pay_for_request(
|
||||
test_key, integration_session, 500, snapshot
|
||||
)
|
||||
results.append(r)
|
||||
|
||||
await integration_session.refresh(test_key)
|
||||
@@ -317,7 +344,11 @@ async def test_child_key_revert_floor_guard(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test that child key reserved_balance also has floor guard on revert."""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
from routstr.auth import (
|
||||
get_reservation_snapshot,
|
||||
pay_for_request,
|
||||
revert_pay_for_request,
|
||||
)
|
||||
|
||||
parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}"
|
||||
child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}"
|
||||
@@ -325,22 +356,26 @@ async def test_child_key_revert_floor_guard(
|
||||
parent_key = ApiKey(
|
||||
hashed_key=parent_key_hash,
|
||||
balance=10000,
|
||||
reserved_balance=500,
|
||||
total_requests=3,
|
||||
reserved_balance=0,
|
||||
total_requests=2,
|
||||
)
|
||||
child_key = ApiKey(
|
||||
hashed_key=child_key_hash,
|
||||
balance=0,
|
||||
reserved_balance=500,
|
||||
total_requests=3,
|
||||
reserved_balance=0,
|
||||
total_requests=2,
|
||||
parent_key_hash=parent_key_hash,
|
||||
)
|
||||
integration_session.add(parent_key)
|
||||
integration_session.add(child_key)
|
||||
await integration_session.commit()
|
||||
await pay_for_request(child_key, 500, integration_session)
|
||||
snapshot = await get_reservation_snapshot(child_key, integration_session)
|
||||
|
||||
# First revert succeeds
|
||||
result1 = await revert_pay_for_request(child_key, integration_session, 500)
|
||||
result1 = await revert_pay_for_request(
|
||||
child_key, integration_session, 500, snapshot
|
||||
)
|
||||
await integration_session.refresh(parent_key)
|
||||
await integration_session.refresh(child_key)
|
||||
|
||||
@@ -349,7 +384,9 @@ async def test_child_key_revert_floor_guard(
|
||||
assert child_key.reserved_balance == 0
|
||||
|
||||
# Second revert is a no-op for both parent and child
|
||||
result2 = await revert_pay_for_request(child_key, integration_session, 500)
|
||||
result2 = await revert_pay_for_request(
|
||||
child_key, integration_session, 500, snapshot
|
||||
)
|
||||
await integration_session.refresh(parent_key)
|
||||
await integration_session.refresh(child_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
|
||||
@@ -20,6 +20,7 @@ from collections.abc import Callable
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from cashu.core.base import MeltQuoteState
|
||||
from httpx import AsyncClient, Response
|
||||
|
||||
from routstr.core.settings import settings
|
||||
@@ -28,7 +29,9 @@ from routstr.core.settings import settings
|
||||
# with the testmint stub that bypasses swapping (see conftest.py).
|
||||
from routstr.wallet import recieve_token as _real_recieve_token
|
||||
|
||||
PRIMARY_MINT = "http://primary:3338"
|
||||
# Match the authenticated fixture's persisted refund mint: existing-key topups
|
||||
# are intentionally constrained to that mint for collateral provenance.
|
||||
PRIMARY_MINT = "http://localhost:3338"
|
||||
|
||||
|
||||
def _make_swap_mocks(
|
||||
@@ -81,7 +84,9 @@ def _make_swap_mocks(
|
||||
quote=f"melt_quote_{invoice}", amount=invoice, fee_reserve=_next_fee()
|
||||
)
|
||||
)
|
||||
mock_token_wallet.melt = AsyncMock(return_value=Mock())
|
||||
mock_token_wallet.melt = AsyncMock(
|
||||
return_value=Mock(state=MeltQuoteState.paid)
|
||||
)
|
||||
|
||||
return mock_token, mock_token_wallet, mock_primary_wallet
|
||||
|
||||
@@ -89,7 +94,12 @@ def _make_swap_mocks(
|
||||
def _wallet_router(primary_wallet: Mock, token_wallet: Mock) -> Callable[..., Mock]:
|
||||
"""Route get_wallet calls to the primary or foreign wallet mock by URL."""
|
||||
|
||||
def fake_get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Mock:
|
||||
def fake_get_wallet(
|
||||
mint_url: str,
|
||||
unit: str = "sat",
|
||||
load: bool = True,
|
||||
**kwargs: object,
|
||||
) -> Mock:
|
||||
return primary_wallet if mint_url == PRIMARY_MINT else token_wallet
|
||||
|
||||
return fake_get_wallet
|
||||
@@ -139,7 +149,7 @@ async def test_topup_retries_when_melt_demands_more_than_quoted(
|
||||
"Mint Error: not enough inputs provided for melt. "
|
||||
"Provided: 179, needed: 180 (Code: 11000)"
|
||||
),
|
||||
Mock(),
|
||||
Mock(state=MeltQuoteState.paid),
|
||||
]
|
||||
|
||||
response = await _post_topup(
|
||||
@@ -174,12 +184,12 @@ async def test_topup_retries_when_quote_fee_exceeds_estimate(
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_returns_400_when_retries_exhausted(
|
||||
async def test_topup_returns_422_when_retries_exhausted(
|
||||
authenticated_client: AsyncClient,
|
||||
) -> None:
|
||||
"""A mint that escalates fee_reserve on every re-quote (1 → 10 → 25 → 50)
|
||||
exhausts the retry budget: clean 400 with an actionable message, melt never
|
||||
executed."""
|
||||
exhausts the retry budget: clean 422 mint_error/too-small taxonomy, melt
|
||||
never executed."""
|
||||
mock_token, token_wallet, primary_wallet = _make_swap_mocks(
|
||||
1000, fee_reserves=[1, 10, 25, 50]
|
||||
)
|
||||
@@ -188,7 +198,7 @@ async def test_topup_returns_400_when_retries_exhausted(
|
||||
authenticated_client, mock_token, token_wallet, primary_wallet
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.status_code == 422
|
||||
assert "too small to cover swap fees" in response.json()["detail"]
|
||||
assert token_wallet.melt_quote.call_count == 4 # estimation + 3 attempts
|
||||
token_wallet.melt.assert_not_called()
|
||||
|
||||
@@ -12,10 +12,7 @@ from sqlmodel import select
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
from .utils import (
|
||||
CashuTokenGenerator,
|
||||
ResponseValidator,
|
||||
)
|
||||
from .utils import ResponseValidator
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@@ -79,29 +76,31 @@ async def test_api_key_generation_invalid_token(
|
||||
# Capture initial state
|
||||
await db_snapshot.capture()
|
||||
|
||||
# Test various invalid tokens
|
||||
invalid_tokens = [
|
||||
CashuTokenGenerator.generate_invalid_token(), # Malformed token
|
||||
"not-a-cashu-token", # Wrong format
|
||||
"cashuA", # Empty token
|
||||
"cashuA" + "x" * 1000, # Invalid base64
|
||||
# Non-Cashu bearer values are invalid API keys (401). Malformed values that
|
||||
# look like Cashu tokens use the shared Cashu taxonomy (400 invalid_token).
|
||||
invalid_tokens: list[tuple[str, int, str | None]] = [
|
||||
("not-a-cashu-token", 401, None),
|
||||
("sk-not-a-real-api-key", 401, None),
|
||||
("cashuA", 400, "invalid_cashu_token"),
|
||||
("cashuA" + "x" * 1000, 400, "invalid_cashu_token"),
|
||||
]
|
||||
|
||||
for invalid_token in invalid_tokens:
|
||||
for invalid_token, expected_status, expected_code in invalid_tokens:
|
||||
integration_client.headers["Authorization"] = f"Bearer {invalid_token}"
|
||||
response = await integration_client.get("/v1/wallet/info")
|
||||
|
||||
# Should fail with 401
|
||||
assert response.status_code == 401, (
|
||||
assert response.status_code == expected_status, (
|
||||
f"Token {invalid_token[:20]}... should be invalid"
|
||||
)
|
||||
|
||||
# Validate error response
|
||||
validator = ResponseValidator()
|
||||
error_validation = validator.validate_error_response(
|
||||
response, expected_status=401, expected_error_key="detail"
|
||||
response, expected_status=expected_status, expected_error_key="detail"
|
||||
)
|
||||
assert error_validation["valid"]
|
||||
if expected_code is not None:
|
||||
assert response.json()["detail"]["error"]["code"] == expected_code
|
||||
|
||||
# Verify no database changes
|
||||
diff = await db_snapshot.diff()
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
"""Restart reconciliation for ambiguous melts, against a real cashu wallet DB.
|
||||
|
||||
The ambiguous-melt path in ``execute_bolt11_payment`` re-reserves proofs with
|
||||
``set_reserved_for_melt(..., quote_id=...)`` after cashu's ``melt()`` clears
|
||||
both the reservation and the ``melt_id`` on a transport error. These tests
|
||||
prove, on cashu's actual sqlite store rather than mocks, that the recovery
|
||||
survives a process restart: a fresh wallet instance on the same database can
|
||||
still find the proofs by ``melt_id`` — the lookup ``get_melt_quote()`` uses to
|
||||
invalidate them on "paid" or release them on "unpaid".
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from cashu.core.base import Proof
|
||||
from cashu.wallet import crud
|
||||
from cashu.wallet.wallet import Wallet
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
QUOTE_ID = "quote-restart-1"
|
||||
|
||||
|
||||
def _proof(secret: str, amount: int = 64) -> Proof:
|
||||
return Proof(
|
||||
id="009a1f293253e41e",
|
||||
amount=amount,
|
||||
secret=secret,
|
||||
C="02bc9097997d81afb2cc7346b5e4345a9346bd2a506eb7958598a72f0cf85163ea",
|
||||
)
|
||||
|
||||
|
||||
async def _wallet(db_dir: Path) -> Wallet:
|
||||
# with_db builds the instance and runs migrations locally; nothing here
|
||||
# talks to a mint.
|
||||
return await Wallet.with_db("https://mint.test", str(db_dir))
|
||||
|
||||
|
||||
async def _seed_ambiguous_melt(wallet: Wallet) -> list[Proof]:
|
||||
"""Reproduce the exact sequence of an ambiguous melt failure.
|
||||
|
||||
1. Proofs exist and are selected for a melt.
|
||||
2. cashu's melt() reserves them with the quote id, then hits a transport
|
||||
error and rolls that back — reservation gone, melt_id gone.
|
||||
3. Our recovery in execute_bolt11_payment re-reserves with the quote id.
|
||||
"""
|
||||
proofs = [_proof("secret-a"), _proof("secret-b", amount=32)]
|
||||
for proof in proofs:
|
||||
await crud.store_proof(proof, db=wallet.db)
|
||||
|
||||
await wallet.set_reserved_for_melt(proofs, reserved=True, quote_id=QUOTE_ID)
|
||||
# cashu's `except` block in melt():
|
||||
await wallet.set_reserved_for_melt(proofs, reserved=False, quote_id=None)
|
||||
# our recovery:
|
||||
await wallet.set_reserved_for_melt(proofs, reserved=True, quote_id=QUOTE_ID)
|
||||
return proofs
|
||||
|
||||
|
||||
async def test_melt_recovery_is_findable_by_quote_after_restart(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
wallet = await _wallet(tmp_path)
|
||||
await _seed_ambiguous_melt(wallet)
|
||||
|
||||
# "Restart": a brand-new wallet on the same database file, as after a
|
||||
# process crash between the melt and any reconciliation.
|
||||
restarted = await _wallet(tmp_path)
|
||||
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
|
||||
|
||||
# This is get_melt_quote()'s own lookup. If it comes back empty, a "paid"
|
||||
# answer can never invalidate these proofs and an "unpaid" answer can
|
||||
# never release them — the strand the send-style re-reserve caused.
|
||||
assert sorted(p.secret for p in found) == ["secret-a", "secret-b"]
|
||||
assert all(p.reserved for p in found)
|
||||
assert all(p.melt_id == QUOTE_ID for p in found)
|
||||
|
||||
|
||||
async def test_send_style_reservation_would_not_be_reconcilable(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""The defect the fix removed, demonstrated on the real store."""
|
||||
wallet = await _wallet(tmp_path)
|
||||
proofs = [_proof("secret-send")]
|
||||
for proof in proofs:
|
||||
await crud.store_proof(proof, db=wallet.db)
|
||||
|
||||
await wallet.set_reserved_for_melt(proofs, reserved=True, quote_id=QUOTE_ID)
|
||||
await wallet.set_reserved_for_melt(proofs, reserved=False, quote_id=None)
|
||||
# The old recovery: reserve as a send, no quote association.
|
||||
await wallet.set_reserved_for_send(proofs, reserved=True)
|
||||
|
||||
restarted = await _wallet(tmp_path)
|
||||
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
|
||||
assert found == [] # reconciliation would never see these proofs
|
||||
|
||||
|
||||
async def test_unpaid_reconciliation_releases_recovered_proofs_after_restart(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""The full recovery arc: crash, restart, mint says unpaid, funds usable."""
|
||||
wallet = await _wallet(tmp_path)
|
||||
await _seed_ambiguous_melt(wallet)
|
||||
|
||||
restarted = await _wallet(tmp_path)
|
||||
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
|
||||
assert len(found) == 2
|
||||
|
||||
# What get_melt_quote() does on an "unpaid" answer.
|
||||
await restarted.set_reserved_for_melt(found, reserved=False, quote_id=None)
|
||||
|
||||
released = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
|
||||
assert released == []
|
||||
all_proofs = await crud.get_proofs(db=restarted.db)
|
||||
assert len(all_proofs) == 2
|
||||
assert all(not p.reserved for p in all_proofs) # spendable again
|
||||
@@ -14,6 +14,7 @@ from httpx import AsyncClient
|
||||
from sqlmodel import select
|
||||
|
||||
from routstr.core.db import ApiKey, CashuTransaction
|
||||
from routstr.wallet import MintConnectionError
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@@ -503,26 +504,19 @@ async def test_mint_unavailability_handling(
|
||||
|
||||
# The global mock in conftest.py is already in place,
|
||||
# so we need to temporarily modify it
|
||||
from unittest.mock import patch
|
||||
raw_error = "Mint unavailable: Connection refused"
|
||||
|
||||
# Make the send_token method raise an exception
|
||||
# Make the send_token method raise a typed mint connection exception.
|
||||
with patch(
|
||||
"routstr.balance.send_token",
|
||||
side_effect=Exception("Mint unavailable: Connection refused"),
|
||||
side_effect=MintConnectionError(raw_error),
|
||||
):
|
||||
# The exception should propagate as a 503 error (Service Unavailable)
|
||||
# But we need to handle it properly
|
||||
try:
|
||||
response = await authenticated_client.post("/v1/wallet/refund")
|
||||
# If we get here, check the status code
|
||||
assert response.status_code == 503
|
||||
assert "Mint service unavailable" in response.json()["detail"]
|
||||
except Exception as e:
|
||||
# If the exception propagates, that's also a failure scenario
|
||||
assert "Mint unavailable" in str(e)
|
||||
response = await authenticated_client.post("/v1/wallet/refund")
|
||||
assert response.status_code == 503
|
||||
assert response.json()["detail"] == "Mint service unavailable"
|
||||
assert raw_error not in response.text
|
||||
|
||||
# Balance should remain unchanged (transaction should roll back)
|
||||
# Note: Current implementation might not handle this perfectly
|
||||
wallet_response = await authenticated_client.get("/v1/wallet/")
|
||||
assert wallet_response.status_code == 200
|
||||
assert wallet_response.json()["balance"] == 10_000_000
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
from contextlib import asynccontextmanager
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.core.admin import get_transactions_api
|
||||
from routstr.core.db import CashuTransaction
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transactions_api_excludes_internal_sweep_claim_timestamp() -> None:
|
||||
transaction = CashuTransaction(
|
||||
token="cashu-token",
|
||||
amount=10,
|
||||
unit="sat",
|
||||
type="out",
|
||||
sweep_started_at=123,
|
||||
)
|
||||
count_result = MagicMock()
|
||||
count_result.one.return_value = 1
|
||||
transactions_result = MagicMock()
|
||||
transactions_result.all.return_value = [transaction]
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(side_effect=[count_result, transactions_result])
|
||||
|
||||
@asynccontextmanager
|
||||
async def create_session(): # type: ignore[no-untyped-def]
|
||||
yield session
|
||||
|
||||
with patch("routstr.core.admin.create_session", create_session):
|
||||
response = await get_transactions_api()
|
||||
|
||||
assert response["total"] == 1
|
||||
assert response["transactions"][0]["token"] == "cashu-token"
|
||||
assert "sweep_started_at" not in response["transactions"][0]
|
||||
@@ -0,0 +1,152 @@
|
||||
import base64
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import routstr.wallet as wallet_module
|
||||
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
|
||||
token = "cashuBoutgoing"
|
||||
send_token = AsyncMock(return_value=token)
|
||||
store_transaction = AsyncMock(return_value=True)
|
||||
|
||||
monkeypatch.setattr(admin, "send_token", send_token)
|
||||
monkeypatch.setattr(admin, "token_mint_url", Mock(return_value=effective_mint))
|
||||
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, "mint_url": effective_mint}
|
||||
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"
|
||||
token = "cashuBrecoverable"
|
||||
|
||||
monkeypatch.setattr(admin, "send_token", AsyncMock(return_value=token))
|
||||
monkeypatch.setattr(admin, "token_mint_url", Mock(return_value=mint))
|
||||
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, "mint_url": mint}
|
||||
critical.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_withdraw_falls_back_from_insufficient_preferred_mint(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
requested_mint = "https://primary.example"
|
||||
actual_mint = "https://secondary.example"
|
||||
proofs = [SimpleNamespace(amount=100, reserved=False, id="00")]
|
||||
token_payload = {
|
||||
"token": [
|
||||
{
|
||||
"mint": actual_mint,
|
||||
"proofs": [
|
||||
{
|
||||
"id": "00",
|
||||
"amount": 75,
|
||||
"secret": "secret",
|
||||
"C": "02" + "00" * 32,
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
"unit": "sat",
|
||||
}
|
||||
token = "cashuA" + base64.urlsafe_b64encode(
|
||||
json.dumps(token_payload).encode()
|
||||
).decode()
|
||||
wallet = SimpleNamespace(
|
||||
keysets={},
|
||||
proofs=proofs,
|
||||
select_to_send=AsyncMock(return_value=(proofs, 0)),
|
||||
serialize_proofs=AsyncMock(return_value=token),
|
||||
set_reserved_for_send=AsyncMock(),
|
||||
)
|
||||
find_funded = AsyncMock(return_value=actual_mint)
|
||||
store_transaction = AsyncMock(return_value=True)
|
||||
|
||||
monkeypatch.setattr(wallet_module, "find_trusted_mint_with_funds", find_funded)
|
||||
monkeypatch.setattr(wallet_module, "get_wallet", AsyncMock(return_value=wallet))
|
||||
monkeypatch.setattr(
|
||||
wallet_module, "get_proofs_per_mint_and_unit", Mock(return_value=proofs)
|
||||
)
|
||||
monkeypatch.setattr(admin, "store_cashu_transaction", store_transaction)
|
||||
|
||||
result = await admin.withdraw(
|
||||
Mock(), admin.WithdrawRequest(amount=75, mint_url=requested_mint)
|
||||
)
|
||||
|
||||
assert result == {"token": token, "mint_url": actual_mint}
|
||||
find_funded.assert_awaited_once_with(
|
||||
75, "sat", requested_mint, force_reload=True
|
||||
)
|
||||
wallet.select_to_send.assert_awaited_once()
|
||||
store_transaction.assert_awaited_once_with(
|
||||
token=token,
|
||||
amount=75,
|
||||
unit="sat",
|
||||
mint_url=actual_mint,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="admin",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_withdraw_maps_true_aggregate_insufficient_funds_to_400(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(
|
||||
admin,
|
||||
"send_token",
|
||||
AsyncMock(
|
||||
side_effect=ValueError(
|
||||
"No trusted mint has 75 sat available; balances={'mint': 0}"
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await admin.withdraw(Mock(), admin.WithdrawRequest(amount=75))
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail == "Insufficient wallet balance"
|
||||
@@ -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]
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
import hashlib
|
||||
from types import SimpleNamespace
|
||||
from typing import AsyncGenerator
|
||||
from typing import AsyncGenerator, cast
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
@@ -12,6 +13,19 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.auth import validate_bearer_key
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.wallet import MintConnectionError
|
||||
|
||||
|
||||
def _value_error_wrapping_transport() -> ValueError:
|
||||
"""A ValueError re-raised ``from`` a real httpx transport error, mirroring
|
||||
``wallet.py`` wrapping a connection failure. The sanitized classifier must
|
||||
still see the mint-unreachable signal through the ``__cause__`` chain."""
|
||||
try:
|
||||
raise httpx.ConnectError("All connection attempts failed")
|
||||
except httpx.ConnectError as exc:
|
||||
err = ValueError("Failed to estimate fees: connection failed")
|
||||
err.__cause__ = exc
|
||||
return err
|
||||
|
||||
|
||||
def _make_engine() -> AsyncEngine:
|
||||
@@ -57,3 +71,271 @@ async def test_failed_first_cashu_redemption_rolls_back_empty_api_key(
|
||||
await validate_bearer_key(token, session)
|
||||
|
||||
assert await session.get(ApiKey, hashed_key) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("error", "expected_status", "expected_type", "expected_message", "expected_code"),
|
||||
[
|
||||
(
|
||||
ValueError("Mint Error: Token already spent. (Code: 11001)"),
|
||||
400,
|
||||
"token_already_spent",
|
||||
"Cashu token already spent",
|
||||
"cashu_token_already_spent",
|
||||
),
|
||||
(
|
||||
# Raw httpx transport error propagated unwrapped from cashu.
|
||||
httpx.ConnectError("All connection attempts failed"),
|
||||
503,
|
||||
"mint_unreachable",
|
||||
"Cashu mint is unreachable",
|
||||
"cashu_mint_unreachable",
|
||||
),
|
||||
(
|
||||
# Typed error raised by wallet.py at a wrap site.
|
||||
MintConnectionError("connect to http://mint:3338 refused"),
|
||||
503,
|
||||
"mint_unreachable",
|
||||
"Cashu mint is unreachable",
|
||||
"cashu_mint_unreachable",
|
||||
),
|
||||
(
|
||||
# ValueError wrapping the httpx error in its __cause__ chain.
|
||||
_value_error_wrapping_transport(),
|
||||
503,
|
||||
"mint_unreachable",
|
||||
"Cashu mint is unreachable",
|
||||
"cashu_mint_unreachable",
|
||||
),
|
||||
(
|
||||
# asyncio.TimeoutError is builtin TimeoutError on 3.11+.
|
||||
TimeoutError("Timed out connecting to Cashu mint http://mint:3338"),
|
||||
503,
|
||||
"mint_unreachable",
|
||||
"Cashu mint is unreachable",
|
||||
"cashu_mint_unreachable",
|
||||
),
|
||||
(
|
||||
ValueError(
|
||||
"Token amount (5 sat) is insufficient to cover melt fees. "
|
||||
"Needed: 7 sat (amount: 5 + fee: 1 + input_fees: 1)"
|
||||
),
|
||||
422,
|
||||
"mint_error",
|
||||
"Token value is too small to cover swap fees",
|
||||
"cashu_token_swap_fees_exceed_amount",
|
||||
),
|
||||
(
|
||||
ValueError(
|
||||
"Failed to estimate fees: Fees (7 sat) exceed token amount (5 sat)"
|
||||
),
|
||||
422,
|
||||
"mint_error",
|
||||
"Token value is too small to cover swap fees",
|
||||
"cashu_token_swap_fees_exceed_amount",
|
||||
),
|
||||
(
|
||||
ValueError(
|
||||
"Failed to melt token from foreign mint http://foreign:3338: boom"
|
||||
),
|
||||
422,
|
||||
"mint_error",
|
||||
"Failed to swap token from foreign mint",
|
||||
"cashu_foreign_mint_swap_failed",
|
||||
),
|
||||
(
|
||||
ValueError("could not decode token"),
|
||||
400,
|
||||
"invalid_token",
|
||||
"Invalid Cashu token",
|
||||
"invalid_cashu_token",
|
||||
),
|
||||
(
|
||||
ValueError("some unexpected wallet condition"),
|
||||
400,
|
||||
"cashu_error",
|
||||
"Failed to redeem Cashu token",
|
||||
"cashu_token_redemption_failed",
|
||||
),
|
||||
(
|
||||
ValueError("Redeemed token amount must be positive, got 0 msats"),
|
||||
400,
|
||||
"cashu_error",
|
||||
"Failed to redeem Cashu token: token yielded no value",
|
||||
"cashu_token_zero_value",
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_redemption_failure_returns_sanitized_error(
|
||||
session: AsyncSession,
|
||||
error: Exception,
|
||||
expected_status: int,
|
||||
expected_type: str,
|
||||
expected_message: str,
|
||||
expected_code: str,
|
||||
) -> None:
|
||||
"""Redemption failures reuse the shared X-Cashu taxonomy (carried in
|
||||
``type``), expose stable sanitized messages and granular machine-readable
|
||||
``code`` values, and leave no orphan ApiKey row."""
|
||||
token = "cashuAredemption_fails_with_specific_error"
|
||||
hashed_key = hashlib.sha256(token.encode()).hexdigest()
|
||||
token_obj = SimpleNamespace(mint="http://mint:3338", unit="sat")
|
||||
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
|
||||
patch(
|
||||
"routstr.auth.credit_balance",
|
||||
new=AsyncMock(side_effect=error),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await validate_bearer_key(token, session)
|
||||
|
||||
assert exc_info.value.status_code == expected_status
|
||||
detail = cast(dict[str, dict[str, object]], exc_info.value.detail)
|
||||
error_detail = detail["error"]
|
||||
assert error_detail["type"] == expected_type
|
||||
assert error_detail["code"] == expected_code
|
||||
assert error_detail["message"] == expected_message
|
||||
assert str(error) not in cast(str, error_detail["message"])
|
||||
assert await session.get(ApiKey, hashed_key) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unexpected_redemption_error_returns_internal_error(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
"""Unexpected (non-wallet) failures surface as generic 500s without
|
||||
leaking internal details, instead of masquerading as token errors."""
|
||||
token = "cashuAredemption_fails_with_internal_error"
|
||||
hashed_key = hashlib.sha256(token.encode()).hexdigest()
|
||||
token_obj = SimpleNamespace(mint="http://mint:3338", unit="sat")
|
||||
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
|
||||
patch(
|
||||
"routstr.auth.credit_balance",
|
||||
new=AsyncMock(side_effect=RuntimeError("db exploded at /var/lib/secret")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await validate_bearer_key(token, session)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
detail = cast(dict[str, dict[str, str]], exc_info.value.detail)
|
||||
error_detail = detail["error"]
|
||||
assert error_detail["code"] == "internal_error"
|
||||
assert "/var/lib/secret" not in error_detail["message"]
|
||||
assert await session.get(ApiKey, hashed_key) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_internal_error_with_invalid_keyword_does_not_masquerade(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
"""A non-wallet fault whose text merely contains "invalid" (but not
|
||||
"token") must fall through to a generic 500, not a 401 token error.
|
||||
|
||||
Guards the anchored `"invalid"/"decode"` + `"token"` gate against stdlib/
|
||||
driver strings like "Invalid isoformat string" leaking as token errors."""
|
||||
token = "cashuAinternal_fault_mentions_invalid"
|
||||
hashed_key = hashlib.sha256(token.encode()).hexdigest()
|
||||
token_obj = SimpleNamespace(mint="http://mint:3338", unit="sat")
|
||||
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
|
||||
patch(
|
||||
"routstr.auth.credit_balance",
|
||||
new=AsyncMock(
|
||||
side_effect=RuntimeError("Invalid isoformat string: '2020-13-99'")
|
||||
),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await validate_bearer_key(token, session)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
detail = cast(dict[str, dict[str, str]], exc_info.value.detail)
|
||||
assert detail["error"]["code"] == "internal_error"
|
||||
assert await session.get(ApiKey, hashed_key) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_primary_msat_token_sets_provenance_without_cashu_mint_duplicate(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
token = "cashuAprimary_msat_token"
|
||||
token_obj = SimpleNamespace(mint="http://primary:3338", unit="msat")
|
||||
credit = AsyncMock(return_value=1_000)
|
||||
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "primary_mint", token_obj.mint),
|
||||
patch.object(settings, "primary_mint_unit", "msat"),
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
|
||||
patch("routstr.auth.credit_balance", new=credit),
|
||||
):
|
||||
key = await validate_bearer_key(token, session)
|
||||
|
||||
assert key.refund_mint_url == token_obj.mint
|
||||
assert key.refund_currency == "msat"
|
||||
credit.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_primary_token_unit_mismatch_is_rejected_before_redemption(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
token = "cashuAprimary_wrong_unit"
|
||||
token_obj = SimpleNamespace(mint="http://primary:3338", unit="sat")
|
||||
credit = AsyncMock(return_value=1_000)
|
||||
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "primary_mint", token_obj.mint),
|
||||
patch.object(settings, "primary_mint_unit", "msat"),
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
|
||||
patch("routstr.auth.credit_balance", new=credit),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await validate_bearer_key(token, session)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
credit.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_malformed_cashu_token_returns_400_invalid_token(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
"""A malformed 'cashu...' token that fails to decode maps to 400
|
||||
invalid_cashu_token (shared taxonomy), not the generic 401 invalid_api_key."""
|
||||
token = "cashuAthis_is_not_a_valid_token"
|
||||
|
||||
with patch(
|
||||
"routstr.auth.deserialize_token_from_string",
|
||||
side_effect=ValueError("unable to decode token: bad base64"),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await validate_bearer_key(token, session)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
detail = cast(dict[str, dict[str, str]], exc_info.value.detail)
|
||||
assert detail["error"]["type"] == "invalid_token"
|
||||
assert detail["error"]["code"] == "invalid_cashu_token"
|
||||
# Raw decoder text must not leak to the client.
|
||||
assert "base64" not in detail["error"]["message"]
|
||||
|
||||
@@ -0,0 +1,736 @@
|
||||
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,
|
||||
_parse_ppq_request_id,
|
||||
_run_auto_topup_cycle,
|
||||
validate_ppq_auto_topup_settings,
|
||||
)
|
||||
from routstr.upstream.ppqai import PPQAIUpstreamProvider
|
||||
from routstr.wallet import Bolt11PaymentAmbiguous, Bolt11PaymentNotAttempted
|
||||
|
||||
|
||||
def test_ppq_claim_parser_rejects_invalid_expiry() -> None:
|
||||
assert (
|
||||
_parse_ppq_request_id("ppq:operation:not-a-timestamp:claimed:invoice:none")
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_balance_rejects_boolean_api_value() -> None:
|
||||
provider = PPQAIUpstreamProvider("secret")
|
||||
provider.check_balance = AsyncMock(return_value={"balance": False}) # type: ignore[method-assign]
|
||||
|
||||
assert await provider.get_balance() is None
|
||||
|
||||
|
||||
def _row() -> MagicMock:
|
||||
row = MagicMock()
|
||||
row.id = "provider-1"
|
||||
row.base_url = "https://provider.test"
|
||||
row.api_key = "secret"
|
||||
row.provider_type = "routstr"
|
||||
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.token_mint_url",
|
||||
return_value="https://fallback-mint.test",
|
||||
),
|
||||
patch("routstr.upstream.auto_topup.create_session", return_value=session),
|
||||
):
|
||||
await _check_and_topup(_row())
|
||||
|
||||
store.assert_awaited_once_with(
|
||||
token="cashu-token",
|
||||
amount=50,
|
||||
unit="sat",
|
||||
mint_url="https://fallback-mint.test",
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
)
|
||||
provider.topup.assert_awaited_once_with("cashu-token")
|
||||
assert transaction.collected is True
|
||||
session.commit.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("outcome", [{"error": "rejected"}, RuntimeError("network")])
|
||||
async def test_auto_topup_failure_leaves_persisted_token_uncollected(
|
||||
outcome: object,
|
||||
) -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=0)
|
||||
provider.topup = AsyncMock(
|
||||
side_effect=outcome if isinstance(outcome, Exception) else None,
|
||||
return_value=outcome,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.send_token",
|
||||
AsyncMock(return_value="cashu-token"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.store_cashu_transaction",
|
||||
AsyncMock(return_value=True),
|
||||
),
|
||||
patch("routstr.upstream.auto_topup.create_session") as create_session,
|
||||
):
|
||||
if isinstance(outcome, Exception):
|
||||
with pytest.raises(RuntimeError):
|
||||
await _check_and_topup(_row())
|
||||
else:
|
||||
await _check_and_topup(_row())
|
||||
|
||||
create_session.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_topup_does_not_send_untracked_token() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=0)
|
||||
provider.topup = AsyncMock()
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.send_token",
|
||||
AsyncMock(return_value="cashu-token"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.store_cashu_transaction",
|
||||
AsyncMock(side_effect=RuntimeError("database unavailable")),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.release_token_reservation",
|
||||
AsyncMock(),
|
||||
) as reclaim,
|
||||
):
|
||||
await _check_and_topup(_row())
|
||||
|
||||
reclaim.assert_awaited_once_with("cashu-token")
|
||||
provider.topup.assert_not_awaited()
|
||||
|
||||
|
||||
def _ppq_row() -> MagicMock:
|
||||
row = MagicMock()
|
||||
row.id = "ppq-provider-1"
|
||||
row.base_url = "https://api.ppq.ai"
|
||||
row.api_key = "secret"
|
||||
row.provider_type = "ppqai"
|
||||
row.provider_settings = json.dumps(
|
||||
{
|
||||
"auto_topup": True,
|
||||
"topup_threshold": 5.0,
|
||||
"topup_amount_limit": 10,
|
||||
}
|
||||
)
|
||||
return row
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_auto_topup_pays_invoice_and_confirms_settlement() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=2.5)
|
||||
provider.initiate_topup = AsyncMock(
|
||||
return_value=MagicMock(
|
||||
invoice_id="invoice-1",
|
||||
payment_request="lnbc-invoice",
|
||||
amount=10,
|
||||
currency="USD",
|
||||
expires_at=None,
|
||||
)
|
||||
)
|
||||
provider.check_topup_status = AsyncMock(return_value=True)
|
||||
plan = MagicMock()
|
||||
plan.invoice_amount_sats = 100
|
||||
plan.maximum_spend_sats = 102
|
||||
plan.quote.amount = 100
|
||||
plan.quote.fee_reserve = 2
|
||||
plan.mint_url = "https://mint-rich.test"
|
||||
plan.unit = "sat"
|
||||
row = _ppq_row()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._claim_ppq_topup",
|
||||
AsyncMock(return_value="operation-1"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.maximum_owner_cashu_balance_sats",
|
||||
AsyncMock(return_value=10_000),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._ppq_spent_last_24h_usd",
|
||||
AsyncMock(return_value=0.0),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.prepare_bolt11_payment",
|
||||
AsyncMock(return_value=plan),
|
||||
) as prepare,
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.execute_bolt11_payment",
|
||||
AsyncMock(return_value=(101, "https://mint-rich.test", "sat")),
|
||||
) as execute,
|
||||
patch("routstr.upstream.auto_topup._record_ppq_invoice", AsyncMock()) as record,
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._record_ppq_payment_spent", AsyncMock()
|
||||
) as record_spent,
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._set_ppq_state_terminal", AsyncMock()
|
||||
) as terminal,
|
||||
patch("routstr.upstream.auto_topup.sats_usd_price", return_value=0.001),
|
||||
):
|
||||
await _check_and_topup(row)
|
||||
|
||||
provider.initiate_topup.assert_awaited_once_with(10)
|
||||
prepare.assert_awaited_once_with("lnbc-invoice")
|
||||
execute.assert_awaited_once_with(plan)
|
||||
record.assert_awaited_once()
|
||||
record_spent.assert_awaited_once_with("operation-1", 101)
|
||||
provider.check_topup_status.assert_awaited_once_with("invoice-1")
|
||||
terminal.assert_awaited_once_with(row, "operation-1", collected=True, swept=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_ambiguous_melt_keeps_claim_and_emits_critical_alert() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=2.5)
|
||||
provider.initiate_topup = AsyncMock(
|
||||
return_value=MagicMock(
|
||||
invoice_id="invoice-1",
|
||||
payment_request="lnbc-invoice",
|
||||
amount=10,
|
||||
currency="USD",
|
||||
expires_at=None,
|
||||
)
|
||||
)
|
||||
plan = MagicMock(maximum_spend_sats=102, mint_url="https://mint.test", unit="sat")
|
||||
plan.quote.amount = 100
|
||||
plan.quote.fee_reserve = 2
|
||||
row = _ppq_row()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._claim_ppq_topup",
|
||||
AsyncMock(return_value="operation-1"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.maximum_owner_cashu_balance_sats",
|
||||
AsyncMock(return_value=10_000),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._ppq_spent_last_24h_usd",
|
||||
AsyncMock(return_value=0.0),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.prepare_bolt11_payment",
|
||||
AsyncMock(return_value=plan),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.execute_bolt11_payment",
|
||||
AsyncMock(side_effect=Bolt11PaymentAmbiguous("ambiguous melt")),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._record_ppq_invoice",
|
||||
AsyncMock(return_value=2_000_000_000),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._mark_ppq_reconcile", AsyncMock()
|
||||
) as reconcile_mark,
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._set_ppq_state_terminal", AsyncMock()
|
||||
) as terminal,
|
||||
patch("routstr.upstream.auto_topup.sats_usd_price", return_value=0.001),
|
||||
patch("routstr.upstream.auto_topup.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(Bolt11PaymentAmbiguous, match="ambiguous melt"):
|
||||
await _check_and_topup(row)
|
||||
|
||||
# The claim is never released — it moves to reconcile for the admin.
|
||||
terminal.assert_not_awaited()
|
||||
reconcile_mark.assert_awaited_once()
|
||||
critical.assert_called_once()
|
||||
assert "admin reconciliation" in critical.call_args.args[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_payment_not_attempted_releases_claim_for_retry() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=2.5)
|
||||
provider.initiate_topup = AsyncMock(
|
||||
return_value=MagicMock(
|
||||
invoice_id="invoice-1",
|
||||
payment_request="lnbc-invoice",
|
||||
amount=10,
|
||||
currency="USD",
|
||||
expires_at=None,
|
||||
)
|
||||
)
|
||||
plan = MagicMock(maximum_spend_sats=102, mint_url="https://mint.test", unit="sat")
|
||||
plan.quote.amount = 100
|
||||
plan.quote.fee_reserve = 2
|
||||
plan.quote.quote = "quote-1"
|
||||
terminal = AsyncMock(return_value=True)
|
||||
row = _ppq_row()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.maximum_owner_cashu_balance_sats",
|
||||
AsyncMock(return_value=10_000),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._ppq_spent_last_24h_usd",
|
||||
AsyncMock(return_value=0.0),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._claim_ppq_topup",
|
||||
AsyncMock(return_value="operation-1"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.prepare_bolt11_payment",
|
||||
AsyncMock(return_value=plan),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._record_ppq_invoice",
|
||||
AsyncMock(return_value=2_000_000_000),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.execute_bolt11_payment",
|
||||
AsyncMock(side_effect=Bolt11PaymentNotAttempted("unpaid")),
|
||||
),
|
||||
patch("routstr.upstream.auto_topup._set_ppq_state_terminal", terminal),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._mark_ppq_reconcile", AsyncMock()
|
||||
) as reconcile,
|
||||
patch("routstr.upstream.auto_topup.sats_usd_price", return_value=0.001),
|
||||
pytest.raises(Bolt11PaymentNotAttempted, match="unpaid"),
|
||||
):
|
||||
await _check_and_topup(row)
|
||||
|
||||
terminal.assert_awaited_once_with(row, "operation-1", collected=False, swept=True)
|
||||
reconcile.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_status_error_after_payment_marks_reconcile_and_alerts() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=2.5)
|
||||
provider.initiate_topup = AsyncMock(
|
||||
return_value=MagicMock(
|
||||
invoice_id="invoice-1",
|
||||
payment_request="lnbc-invoice",
|
||||
amount=10,
|
||||
currency="USD",
|
||||
expires_at=None,
|
||||
)
|
||||
)
|
||||
provider.check_topup_status = AsyncMock(side_effect=RuntimeError("PPQ 502"))
|
||||
plan = MagicMock(maximum_spend_sats=102, mint_url="https://mint.test", unit="sat")
|
||||
plan.quote.amount = 100
|
||||
plan.quote.fee_reserve = 2
|
||||
plan.quote.quote = "quote-1"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.maximum_owner_cashu_balance_sats",
|
||||
AsyncMock(return_value=10_000),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._ppq_spent_last_24h_usd",
|
||||
AsyncMock(return_value=0.0),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._claim_ppq_topup",
|
||||
AsyncMock(return_value="operation-1"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.prepare_bolt11_payment",
|
||||
AsyncMock(return_value=plan),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._record_ppq_invoice",
|
||||
AsyncMock(return_value=2_000_000_000),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.execute_bolt11_payment",
|
||||
AsyncMock(return_value=(101, "https://mint.test", "sat")),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._record_ppq_payment_spent", AsyncMock()
|
||||
) as spent,
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._mark_ppq_reconcile", AsyncMock()
|
||||
) as reconcile,
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._set_ppq_state_terminal", AsyncMock()
|
||||
) as terminal,
|
||||
patch("routstr.upstream.auto_topup.sats_usd_price", return_value=0.001),
|
||||
patch("routstr.upstream.auto_topup.logger.critical") as critical,
|
||||
):
|
||||
await _check_and_topup(_ppq_row())
|
||||
|
||||
spent.assert_awaited_once_with("operation-1", 101)
|
||||
reconcile.assert_awaited_once()
|
||||
terminal.assert_not_awaited()
|
||||
assert "settlement polling failed" in critical.call_args.args[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_preflight_funding_check_happens_before_invoice_creation() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=2.5)
|
||||
provider.initiate_topup = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.maximum_owner_cashu_balance_sats",
|
||||
AsyncMock(return_value=1),
|
||||
),
|
||||
patch("routstr.upstream.auto_topup._claim_ppq_topup", AsyncMock()) as claim,
|
||||
patch("routstr.upstream.auto_topup.sats_usd_price", return_value=0.001),
|
||||
):
|
||||
await _check_and_topup(_ppq_row())
|
||||
|
||||
provider.initiate_topup.assert_not_awaited()
|
||||
claim.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_active_claim_at_cycle_start_suppresses_topup_for_whole_cycle() -> None:
|
||||
row = _ppq_row()
|
||||
row.id = 1
|
||||
session = AsyncMock()
|
||||
result = MagicMock()
|
||||
result.all.return_value = [row]
|
||||
session.exec.return_value = result
|
||||
context = MagicMock()
|
||||
context.__aenter__ = AsyncMock(return_value=session)
|
||||
context.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_all_ppq_claims",
|
||||
AsyncMock(return_value={1}),
|
||||
),
|
||||
patch("routstr.upstream.auto_topup.create_session", return_value=context),
|
||||
patch("routstr.upstream.auto_topup._check_and_topup", AsyncMock()) as check,
|
||||
):
|
||||
await _run_auto_topup_cycle()
|
||||
|
||||
check.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_auto_topup_skips_when_balance_meets_threshold() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=5.0)
|
||||
provider.initiate_topup = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
):
|
||||
await _check_and_topup(_ppq_row())
|
||||
|
||||
provider.initiate_topup.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_auto_topup_skips_when_daily_spend_cap_reached() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=2.5)
|
||||
provider.initiate_topup = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.maximum_owner_cashu_balance_sats",
|
||||
AsyncMock(return_value=10_000_000),
|
||||
),
|
||||
# 1000 USD already spent, exactly the daily cap: the next 10 USD
|
||||
# top-up must be refused.
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._ppq_spent_last_24h_usd",
|
||||
AsyncMock(return_value=1000.0),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._claim_ppq_topup",
|
||||
AsyncMock(),
|
||||
) as claim,
|
||||
patch("routstr.upstream.auto_topup.sats_usd_price", return_value=0.001),
|
||||
):
|
||||
await _check_and_topup(_ppq_row())
|
||||
|
||||
claim.assert_not_awaited()
|
||||
provider.initiate_topup.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_pending_attempt_suppresses_duplicate_topup() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=True),
|
||||
),
|
||||
):
|
||||
await _check_and_topup(_ppq_row())
|
||||
|
||||
provider.get_balance.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_auto_topup_rejects_non_finite_balance() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=float("nan"))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch("routstr.upstream.auto_topup._claim_ppq_topup", AsyncMock()) as claim,
|
||||
):
|
||||
await _check_and_topup(_ppq_row())
|
||||
|
||||
claim.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_settled_topup_alerts_when_its_claim_was_already_released() -> None:
|
||||
provider = MagicMock()
|
||||
provider.get_balance = AsyncMock(return_value=2.5)
|
||||
provider.initiate_topup = AsyncMock(
|
||||
return_value=MagicMock(
|
||||
invoice_id="invoice-1",
|
||||
payment_request="lnbc-invoice",
|
||||
amount=10,
|
||||
currency="USD",
|
||||
expires_at=None,
|
||||
)
|
||||
)
|
||||
provider.check_topup_status = AsyncMock(return_value=True)
|
||||
plan = MagicMock()
|
||||
plan.maximum_spend_sats = 102
|
||||
plan.quote.amount = 100
|
||||
plan.quote.fee_reserve = 2
|
||||
plan.mint_url = "https://mint-rich.test"
|
||||
plan.unit = "sat"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row",
|
||||
return_value=provider,
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._reconcile_ppq_state",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._claim_ppq_topup",
|
||||
AsyncMock(return_value="operation-1"),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.prepare_bolt11_payment",
|
||||
AsyncMock(return_value=plan),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.maximum_owner_cashu_balance_sats",
|
||||
AsyncMock(return_value=10_000),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._ppq_spent_last_24h_usd",
|
||||
AsyncMock(return_value=0.0),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.execute_bolt11_payment",
|
||||
AsyncMock(return_value=(101, "https://mint-rich.test", "sat")),
|
||||
),
|
||||
patch("routstr.upstream.auto_topup._record_ppq_invoice", AsyncMock()),
|
||||
patch("routstr.upstream.auto_topup._record_ppq_payment_spent", AsyncMock()),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup._set_ppq_state_terminal",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch("routstr.upstream.auto_topup.sats_usd_price", return_value=0.001),
|
||||
patch("routstr.upstream.auto_topup.logger") as log,
|
||||
):
|
||||
await _check_and_topup(_ppq_row())
|
||||
|
||||
assert any(
|
||||
"claim was already released" in call.args[0]
|
||||
for call in log.critical.call_args_list
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("settings", "expected"),
|
||||
[
|
||||
({"auto_topup": False, "topup_threshold": -1}, None),
|
||||
(
|
||||
{"auto_topup": True, "topup_threshold": 5, "topup_amount_limit": 10},
|
||||
None,
|
||||
),
|
||||
(
|
||||
{"auto_topup": True, "topup_threshold": None, "topup_amount_limit": 10},
|
||||
"threshold",
|
||||
),
|
||||
(
|
||||
{"auto_topup": True, "topup_threshold": 5, "topup_amount_limit": 0.5},
|
||||
"whole number",
|
||||
),
|
||||
(
|
||||
{"auto_topup": True, "topup_threshold": 5, "topup_amount_limit": 5000},
|
||||
"between",
|
||||
),
|
||||
(
|
||||
{"auto_topup": True, "topup_threshold": True, "topup_amount_limit": 10},
|
||||
"threshold",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_ppq_auto_topup_settings_validation(
|
||||
settings: dict, expected: str | None
|
||||
) -> None:
|
||||
problem = validate_ppq_auto_topup_settings(settings)
|
||||
if expected is None:
|
||||
assert problem is None
|
||||
else:
|
||||
assert problem is not None and expected in problem
|
||||
|
||||
|
||||
def test_ppq_auto_topup_settings_validation_survives_huge_json_integers() -> None:
|
||||
# json.loads happily produces integers past float range; float() raises
|
||||
# OverflowError there instead of returning inf.
|
||||
problem = validate_ppq_auto_topup_settings(
|
||||
{"auto_topup": True, "topup_threshold": 10**400, "topup_amount_limit": 10}
|
||||
)
|
||||
assert problem is not None and "threshold" in problem
|
||||
+448
-3
@@ -1,12 +1,13 @@
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from routstr.balance import refund_wallet_endpoint
|
||||
from routstr.balance import refund_wallet_endpoint, topup_wallet_endpoint
|
||||
from routstr.core.db import ApiKey, CashuTransaction
|
||||
from routstr.wallet import credit_balance
|
||||
from routstr.wallet import MintConnectionError, credit_balance
|
||||
|
||||
|
||||
def _make_cashu_tx(
|
||||
@@ -103,6 +104,64 @@ async def test_refund_x_cashu_not_found_raises_404() -> None:
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_x_cashu_pending_raises_425() -> None:
|
||||
"""in row exists with a request_id but out row not yet created → 425.
|
||||
|
||||
This is the race condition where /v1/wallet/refund is polled while the
|
||||
upstream request is still in flight. The endpoint must signal "retry"
|
||||
rather than a permanent 404.
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
x_cashu_token = "cashuApending_token"
|
||||
in_tx = _make_cashu_tx(
|
||||
token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-pending"
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(None)])
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-somekey",
|
||||
x_cashu=x_cashu_token,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 425
|
||||
assert exc_info.value.headers == {"Retry-After": "2"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_x_cashu_in_tx_without_request_id_raises_404() -> None:
|
||||
"""in row exists but has no request_id (cannot link to a refund) → 404.
|
||||
|
||||
This is a genuine "no refund will ever exist" case, distinct from the
|
||||
pending 425 path.
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
x_cashu_token = "cashuAnoreqid_token"
|
||||
in_tx = _make_cashu_tx(
|
||||
token=x_cashu_token, amount=0, unit="msat", type="in", request_id=None
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(side_effect=[_exec_result(in_tx)])
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-somekey",
|
||||
x_cashu=x_cashu_token,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_x_cashu_swept_raises_410() -> None:
|
||||
from fastapi import HTTPException
|
||||
@@ -162,6 +221,78 @@ def _make_api_key(
|
||||
return key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_returns_persisted_token_after_cache_loss() -> None:
|
||||
key = _make_api_key(balance=0, refund_currency="sat")
|
||||
refund_token = "cashuApersisted_refund_token"
|
||||
refund_tx = _make_cashu_tx(
|
||||
token=refund_token,
|
||||
amount=5,
|
||||
unit="sat",
|
||||
type="out",
|
||||
request_id=None,
|
||||
)
|
||||
refund_tx.source = "apikey"
|
||||
refund_tx.api_key_hashed_key = key.hashed_key
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=_exec_result(refund_tx))
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance.send_token", AsyncMock()) as mock_send_token,
|
||||
):
|
||||
result = await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert result == {"token": refund_token, "sats": "5"}
|
||||
assert refund_tx.collected is True
|
||||
session.add.assert_called_once_with(refund_tx)
|
||||
session.commit.assert_awaited_once()
|
||||
mock_send_token.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_rejects_persisted_token_after_sweep() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=0, refund_currency="sat")
|
||||
refund_tx = _make_cashu_tx(
|
||||
token="cashuAswept_apikey_refund",
|
||||
amount=5,
|
||||
unit="sat",
|
||||
request_id=None,
|
||||
swept=True,
|
||||
)
|
||||
refund_tx.source = "apikey"
|
||||
refund_tx.api_key_hashed_key = key.hashed_key
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=_exec_result(refund_tx))
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 410
|
||||
assert exc_info.value.detail == "Refund has been swept"
|
||||
session.add.assert_not_called()
|
||||
session.commit.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> None:
|
||||
key = _make_api_key(balance=5000, refund_currency="sat")
|
||||
@@ -339,7 +470,10 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(side_effect=Exception("mint down"))),
|
||||
patch(
|
||||
"routstr.balance.send_token",
|
||||
AsyncMock(side_effect=MintConnectionError("raw mint outage detail")),
|
||||
),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
@@ -353,10 +487,46 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 503
|
||||
assert exc_info.value.detail == "Mint service unavailable"
|
||||
assert "raw mint outage detail" not in exc_info.value.detail
|
||||
# Verify two exec calls: debit + restore
|
||||
assert session.exec.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_generic_failure_is_sanitized_500() -> None:
|
||||
"""Unexpected send-side failures restore balance without leaking exception text."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=5000, refund_currency="sat")
|
||||
raw_error = "database secret token raw-mint-response"
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)])
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(side_effect=RuntimeError(raw_error))),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
patch("routstr.balance.logger"),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert exc_info.value.detail == "Refund failed"
|
||||
assert raw_error not in exc_info.value.detail
|
||||
assert session.exec.await_count == 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# no-create guarantee: fresh Cashu/unknown sk- tokens must not create API keys
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -401,3 +571,278 @@ async def test_refund_unknown_sk_bearer_returns_401() -> None:
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
session.get.assert_awaited_once()
|
||||
|
||||
|
||||
# --- Topup redemption error taxonomy (POST /v1/wallet/topup) ------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"error",
|
||||
[
|
||||
httpx.ConnectError("All connection attempts failed"),
|
||||
MintConnectionError("connect to mint refused"),
|
||||
TimeoutError("timed out connecting to mint"),
|
||||
],
|
||||
)
|
||||
async def test_topup_mint_unreachable_returns_503(error: Exception) -> None:
|
||||
"""A down mint must surface 503 (retryable), not 400 or 500 — the token is
|
||||
fine, so the client should retry once the mint recovers."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await topup_wallet_endpoint(
|
||||
cashu_token="cashuAtoken", key=key, session=session
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 503
|
||||
assert exc_info.value.detail == "Cashu mint is unreachable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.wallet import SourceMintConnectionError
|
||||
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
error = SourceMintConnectionError("Issuing Cashu mint is unreachable")
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await topup_wallet_endpoint(
|
||||
cashu_token="cashuAtoken", key=key, session=session
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 503
|
||||
assert "cannot be redeemed at another mint" in exc_info.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_already_spent_still_returns_400() -> None:
|
||||
"""Regression: the mint-unreachable short-circuit must not swallow the
|
||||
existing ValueError substring buckets."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch(
|
||||
"routstr.balance.credit_balance",
|
||||
AsyncMock(side_effect=ValueError("Token already spent")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await topup_wallet_endpoint(
|
||||
cashu_token="cashuAtoken", key=key, session=session
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail == "Cashu token already spent"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_zero_value_returns_400_zero_value_message() -> None:
|
||||
"""A dust/zero redemption maps to the documented zero-value message, not the
|
||||
generic redemption-failed one."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch(
|
||||
"routstr.balance.credit_balance",
|
||||
AsyncMock(
|
||||
side_effect=ValueError("Redeemed token amount must be positive, got 0 msats")
|
||||
),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await topup_wallet_endpoint(
|
||||
cashu_token="cashuAtoken", key=key, session=session
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail == "Failed to redeem Cashu token: token yielded no value"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_token_consumed_returns_500() -> None:
|
||||
"""A post-redemption crediting failure (token spent) is a non-retryable 500,
|
||||
not a 4xx that invites a retry."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.wallet import TokenConsumedError
|
||||
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch(
|
||||
"routstr.balance.credit_balance",
|
||||
AsyncMock(side_effect=TokenConsumedError("credit failed")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await topup_wallet_endpoint(
|
||||
cashu_token="cashuAtoken", key=key, session=session
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert exc_info.value.detail == (
|
||||
"Token was redeemed but could not be credited; do not retry"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("error", "expected_status", "expected_detail"),
|
||||
[
|
||||
(
|
||||
ValueError(
|
||||
"Failed to estimate fees: Fees (7 sat) exceed token amount (5 sat)"
|
||||
),
|
||||
422,
|
||||
"Token value is too small to cover swap fees",
|
||||
),
|
||||
(
|
||||
ValueError(
|
||||
"Token amount (5 sat) is insufficient to cover melt fees."
|
||||
),
|
||||
422,
|
||||
"Token value is too small to cover swap fees",
|
||||
),
|
||||
(
|
||||
ValueError("Failed to melt token from foreign mint http://m: boom"),
|
||||
422,
|
||||
"Failed to swap token from foreign mint",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_topup_fee_and_swap_failures_return_422(
|
||||
error: Exception, expected_status: int, expected_detail: str
|
||||
) -> None:
|
||||
"""Fee/swap failures map to 422 (shared taxonomy), matching the bearer and
|
||||
X-Cashu paths — previously top-up flattened these to 400."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await topup_wallet_endpoint(
|
||||
cashu_token="cashuAtoken", key=key, session=session
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == expected_status
|
||||
assert exc_info.value.detail == expected_detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_unexpected_non_valueerror_returns_500() -> None:
|
||||
"""A non-ValueError, non-transport fault is an internal error (500), not a
|
||||
sanitized 400 — the merged except must preserve this."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch(
|
||||
"routstr.balance.credit_balance",
|
||||
AsyncMock(side_effect=RuntimeError("db exploded")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await topup_wallet_endpoint(
|
||||
cashu_token="cashuAtoken", key=key, session=session
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert exc_info.value.detail == "Internal server error"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_ambiguous_melt_does_not_restore_balance() -> None:
|
||||
"""An ambiguous LNURL melt may still settle: the debit must be kept."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.payment.lnurl import MeltOutcomeAmbiguousError
|
||||
|
||||
key = _make_api_key(balance=5000, refund_address="user@ln.example.com")
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=MagicMock(rowcount=1))
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
patch(
|
||||
"routstr.balance.send_to_lnurl",
|
||||
AsyncMock(side_effect=MeltOutcomeAmbiguousError("outcome is ambiguous")),
|
||||
),
|
||||
patch("routstr.balance._restore_balance", AsyncMock()) as mock_restore,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 502
|
||||
mock_restore.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_clean_failure_still_restores_balance() -> None:
|
||||
"""A definitively failed melt must keep restoring the debited balance."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=5000, refund_address="user@ln.example.com")
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=MagicMock(rowcount=1))
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
patch(
|
||||
"routstr.balance.send_to_lnurl",
|
||||
AsyncMock(side_effect=RuntimeError("mint rejected melt")),
|
||||
),
|
||||
patch("routstr.balance._restore_balance", AsyncMock()) as mock_restore,
|
||||
):
|
||||
with pytest.raises(HTTPException):
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
mock_restore.assert_awaited_once()
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from httpx import ASGITransport, AsyncClient
|
||||
|
||||
from routstr import balance as balance_module
|
||||
from routstr.core.db import get_session
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_balance_accepts_large_cashu_token_in_post_body(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
token = "cashuA" + "x" * 20_000
|
||||
key = SimpleNamespace(hashed_key="hashed", balance=123_000)
|
||||
validate_bearer_key = AsyncMock(return_value=key)
|
||||
session = AsyncMock()
|
||||
monkeypatch.setattr(balance_module, "validate_bearer_key", validate_bearer_key)
|
||||
|
||||
async def override_get_session(): # type: ignore[no-untyped-def]
|
||||
yield session
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(balance_module.balance_router)
|
||||
app.dependency_overrides[get_session] = override_get_session
|
||||
|
||||
async with AsyncClient(
|
||||
transport=ASGITransport(app=app), # type: ignore[arg-type]
|
||||
base_url="http://test",
|
||||
) as client:
|
||||
response = await client.post(
|
||||
"/v1/balance/create",
|
||||
json={"initial_balance_token": token},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"api_key": "sk-hashed", "balance": 123_000}
|
||||
validate_bearer_key.assert_awaited_once_with(token, session)
|
||||
@@ -0,0 +1,115 @@
|
||||
"""Real-DB coverage for db.balances_by_mint_and_unit.
|
||||
|
||||
Verifies the grouped liability query used by fetch_all_balances: it sums
|
||||
balances per (mint_url, unit), filters to the requested mints/units, excludes
|
||||
NULL mint/currency rows, and returns nothing for empty inputs.
|
||||
"""
|
||||
|
||||
from typing import AsyncGenerator
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
from sqlalchemy.pool import StaticPool
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import (
|
||||
ApiKey,
|
||||
balance_for_mint_and_unit,
|
||||
balances_by_mint_and_unit,
|
||||
)
|
||||
|
||||
|
||||
def _make_engine() -> AsyncEngine:
|
||||
return create_async_engine(
|
||||
"sqlite+aiosqlite://",
|
||||
poolclass=StaticPool,
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def session() -> "AsyncGenerator[AsyncSession, None]":
|
||||
engine = _make_engine()
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
db_session = AsyncSession(engine, expire_on_commit=False)
|
||||
try:
|
||||
yield db_session
|
||||
finally:
|
||||
await db_session.close()
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def _add_key(
|
||||
session: AsyncSession,
|
||||
hashed_key: str,
|
||||
balance: int,
|
||||
mint_url: str | None,
|
||||
currency: str | None,
|
||||
) -> None:
|
||||
session.add(
|
||||
ApiKey(
|
||||
hashed_key=hashed_key,
|
||||
balance=balance,
|
||||
refund_mint_url=mint_url,
|
||||
refund_currency=currency,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sums_and_groups_by_mint_and_unit(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 500, "http://m1", "sat")
|
||||
await _add_key(session, "c", 7000, "http://m1", "msat")
|
||||
await _add_key(session, "d", 200, "http://m2", "sat")
|
||||
|
||||
result = await balances_by_mint_and_unit(
|
||||
session, ["http://m1", "http://m2"], ["sat", "msat"]
|
||||
)
|
||||
|
||||
assert result[("http://m1", "sat")] == 1500
|
||||
assert result[("http://m1", "msat")] == 7000
|
||||
assert result[("http://m2", "sat")] == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_filters_out_unrequested_mints_and_units(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://wanted", "sat")
|
||||
await _add_key(session, "b", 999, "http://other", "sat")
|
||||
await _add_key(session, "c", 888, "http://wanted", "usd")
|
||||
|
||||
result = await balances_by_mint_and_unit(session, ["http://wanted"], ["sat"])
|
||||
|
||||
assert result == {("http://wanted", "sat"): 1000}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_excludes_rows_with_null_mint_or_currency(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 4242, None, None)
|
||||
|
||||
result = await balances_by_mint_and_unit(session, ["http://m1"], ["sat"])
|
||||
|
||||
assert result == {("http://m1", "sat"): 1000}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scalar_balance_for_one_mint_and_unit(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 500, "http://m1", "sat")
|
||||
await _add_key(session, "c", 9000, "http://m1", "msat")
|
||||
await _add_key(session, "d", 700, "http://m2", "sat")
|
||||
|
||||
assert await balance_for_mint_and_unit(session, "http://m1", "sat") == 1500
|
||||
assert await balance_for_mint_and_unit(session, "http://missing", "sat") == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_inputs_return_empty_mapping(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
|
||||
assert await balances_by_mint_and_unit(session, [], ["sat"]) == {}
|
||||
assert await balances_by_mint_and_unit(session, ["http://m1"], []) == {}
|
||||
@@ -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()
|
||||
@@ -15,6 +15,7 @@ os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
|
||||
|
||||
from routstr.core.settings import settings
|
||||
from routstr.payment.cost_calculation import CostData, MaxCostData, calculate_cost
|
||||
from routstr.payment.models import Architecture, Model, Pricing
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
@@ -465,9 +466,306 @@ 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 == 995
|
||||
assert result.output_msats == 3476
|
||||
assert result.cache_read_msats == 758
|
||||
assert result.cache_creation_msats == 0
|
||||
assert result.input_msats + result.output_msats == result.total_msats == 4471
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usd_cache_breakdown_matches_token_priced_path(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Authoritative USD totals must retain model-specific cache-rate ratios."""
|
||||
monkeypatch.setattr(settings, "fixed_pricing", False)
|
||||
model = Model(
|
||||
id="cache-priced-model",
|
||||
name="cache-priced-model",
|
||||
created=0,
|
||||
description="",
|
||||
context_length=8192,
|
||||
architecture=Architecture(
|
||||
modality="text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="test",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(prompt=0.01, completion=0.02),
|
||||
sats_pricing=Pricing(
|
||||
prompt=0.01,
|
||||
completion=0.02,
|
||||
input_cache_read=0.001,
|
||||
input_cache_write=0.01,
|
||||
),
|
||||
per_request_limits=None,
|
||||
top_provider=None,
|
||||
)
|
||||
usage = {
|
||||
"prompt_tokens": 1000,
|
||||
"completion_tokens": 100,
|
||||
"prompt_tokens_details": {"cached_tokens": 900},
|
||||
}
|
||||
|
||||
token_result = await calculate_cost(
|
||||
{"model": model.id, "usage": usage},
|
||||
max_cost=100_000,
|
||||
model_obj=model,
|
||||
)
|
||||
usd_result = await calculate_cost(
|
||||
{
|
||||
"model": model.id,
|
||||
"usage": {
|
||||
**usage,
|
||||
"cost": 0.000195,
|
||||
"cost_details": {
|
||||
"input_cost": 0.000095,
|
||||
"output_cost": 0.0001,
|
||||
},
|
||||
},
|
||||
},
|
||||
max_cost=100_000,
|
||||
model_obj=model,
|
||||
provider_fee=1.0,
|
||||
)
|
||||
|
||||
assert isinstance(token_result, CostData)
|
||||
assert isinstance(usd_result, CostData)
|
||||
assert usd_result.total_msats == token_result.total_msats == 3900
|
||||
assert usd_result.input_msats + usd_result.output_msats == usd_result.total_msats
|
||||
assert usd_result.cache_read_msats == token_result.cache_read_msats == 900
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_usd_cache_breakdown_does_not_absorb_total_rounding_remainder(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Sub-msat cache components truncate like the token-priced path."""
|
||||
monkeypatch.setattr(settings, "fixed_pricing", False)
|
||||
model = Model(
|
||||
id="sub-msat-cache-model",
|
||||
name="sub-msat-cache-model",
|
||||
created=0,
|
||||
description="",
|
||||
context_length=8192,
|
||||
architecture=Architecture(
|
||||
modality="text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="test",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(prompt=0.001, completion=0.001),
|
||||
sats_pricing=Pricing(
|
||||
prompt=0.001,
|
||||
completion=0.001,
|
||||
input_cache_write=0.0006,
|
||||
),
|
||||
per_request_limits=None,
|
||||
top_provider=None,
|
||||
)
|
||||
usage = {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
"cache_creation_input_tokens": 1,
|
||||
}
|
||||
|
||||
token_result = await calculate_cost(
|
||||
{"model": model.id, "usage": usage},
|
||||
max_cost=100_000,
|
||||
model_obj=model,
|
||||
)
|
||||
usd_result = await calculate_cost(
|
||||
{
|
||||
"model": model.id,
|
||||
"usage": {
|
||||
**usage,
|
||||
"cost": 0.00000003,
|
||||
"cost_details": {"input_cost": 0.00000003},
|
||||
},
|
||||
},
|
||||
max_cost=100_000,
|
||||
model_obj=model,
|
||||
provider_fee=1.0,
|
||||
)
|
||||
|
||||
assert isinstance(token_result, CostData)
|
||||
assert isinstance(usd_result, CostData)
|
||||
assert usd_result.total_msats == token_result.total_msats == 1
|
||||
assert usd_result.cache_creation_msats == token_result.cache_creation_msats == 0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# 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 == 926547
|
||||
assert result.output_msats == 13727
|
||||
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.cache_read_msats == 897966
|
||||
assert result.cache_creation_msats == 0
|
||||
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,113 @@
|
||||
"""Response-contract tests for Routstr cost metadata across paid paths."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||
|
||||
from routstr.core.db import ApiKey # noqa: E402
|
||||
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
|
||||
|
||||
COST_DATA = {
|
||||
"base_msats": 0,
|
||||
"input_msats": 1_200,
|
||||
"output_msats": 300,
|
||||
"total_msats": 1_500,
|
||||
"total_usd": 0.0001,
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 3,
|
||||
"cache_read_input_tokens": 8,
|
||||
"cache_creation_input_tokens": 2,
|
||||
"cache_read_msats": 80,
|
||||
"cache_creation_msats": 40,
|
||||
}
|
||||
|
||||
|
||||
def _provider() -> BaseUpstreamProvider:
|
||||
return BaseUpstreamProvider(base_url="http://test", api_key="upstream-key")
|
||||
|
||||
|
||||
def _key() -> ApiKey:
|
||||
return ApiKey(hashed_key="abcdef0123" * 4, balance=1_000_000)
|
||||
|
||||
|
||||
def _session() -> Any:
|
||||
session = MagicMock()
|
||||
session.refresh = AsyncMock()
|
||||
return session
|
||||
|
||||
|
||||
def _upstream_response(payload: dict) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
200,
|
||||
json=payload,
|
||||
request=httpx.Request("POST", "http://test"),
|
||||
)
|
||||
|
||||
|
||||
def _assert_cost_contract(response: Any) -> None:
|
||||
body = json.loads(response.body)
|
||||
assert body["usage"]["cost"] == {
|
||||
"base_msats": 0,
|
||||
"input_msats": 1_200,
|
||||
"output_msats": 300,
|
||||
"total_msats": 1_500,
|
||||
"total_usd": 0.0001,
|
||||
"cache_read_input_tokens": 8,
|
||||
"cache_creation_input_tokens": 2,
|
||||
"cache_read_msats": 80,
|
||||
"cache_creation_msats": 40,
|
||||
}
|
||||
assert response.headers["X-Routstr-Cost-Msats"] == "1500"
|
||||
assert response.headers["X-Routstr-Input-Cost-Msats"] == "1200"
|
||||
assert response.headers["X-Routstr-Output-Cost-Msats"] == "300"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_chat_completion_uses_shared_cost_contract() -> None:
|
||||
provider = _provider()
|
||||
with patch(
|
||||
"routstr.upstream.base.adjust_payment_for_tokens",
|
||||
new=AsyncMock(return_value=dict(COST_DATA)),
|
||||
):
|
||||
response = await provider.handle_non_streaming_chat_completion(
|
||||
_upstream_response(
|
||||
{
|
||||
"model": "test-model",
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 3},
|
||||
}
|
||||
),
|
||||
_key(),
|
||||
_session(),
|
||||
deducted_max_cost=10_000,
|
||||
)
|
||||
|
||||
_assert_cost_contract(response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_responses_completion_uses_shared_cost_contract() -> None:
|
||||
provider = _provider()
|
||||
with patch(
|
||||
"routstr.upstream.base.adjust_payment_for_tokens",
|
||||
new=AsyncMock(return_value=dict(COST_DATA)),
|
||||
):
|
||||
response = await provider.handle_non_streaming_responses_completion(
|
||||
_upstream_response(
|
||||
{
|
||||
"model": "test-model",
|
||||
"usage": {"input_tokens": 10, "output_tokens": 3},
|
||||
}
|
||||
),
|
||||
_key(),
|
||||
_session(),
|
||||
deducted_max_cost=10_000,
|
||||
)
|
||||
|
||||
_assert_cost_contract(response)
|
||||
@@ -0,0 +1,133 @@
|
||||
"""Coverage tests for admin.py (currently 35%).
|
||||
|
||||
Tests admin endpoints that are testable without full app setup:
|
||||
withdraw validation, authentication guards, and slug validation.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, 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.send_token",
|
||||
new=AsyncMock(
|
||||
side_effect=ValueError(
|
||||
"No trusted mint has 1000000 sat available; balances={}"
|
||||
)
|
||||
),
|
||||
):
|
||||
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)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 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,217 @@
|
||||
"""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")
|
||||
await p.on_upstream_error_redirect(402, "Insufficient balance")
|
||||
|
||||
|
||||
@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")
|
||||
await p.on_upstream_error_redirect(429, "Rate limited")
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# _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: dict[str, str] = {}
|
||||
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,85 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from routstr.core import db
|
||||
from routstr.core.db import create_db_engine
|
||||
from routstr.core.settings import settings
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_engine_uses_validated_bounded_pool_settings(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: object
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_size", 12)
|
||||
monkeypatch.setattr(settings, "database_max_overflow", 3)
|
||||
monkeypatch.setattr(settings, "database_pool_timeout", 2.5)
|
||||
monkeypatch.setattr(settings, "database_pool_recycle", 900)
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
|
||||
engine = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/pool.db")
|
||||
try:
|
||||
assert engine.pool.size() == 12 # type: ignore[attr-defined]
|
||||
assert engine.pool._max_overflow == 3 # type: ignore[attr-defined]
|
||||
assert engine.pool._timeout == 2.5 # type: ignore[attr-defined]
|
||||
assert engine.pool._recycle == 900
|
||||
assert engine.pool._pre_ping is False
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_memory_sqlite_keeps_static_pool(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", True)
|
||||
engine = create_db_engine("sqlite+aiosqlite://")
|
||||
try:
|
||||
assert isinstance(engine.pool, StaticPool)
|
||||
assert engine.pool._pre_ping is True
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def test_non_sqlite_backend_enables_pre_ping_automatically(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
fake_engine = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(db, "create_async_engine", return_value=fake_engine) as factory,
|
||||
patch.object(db.event, "listen") as listen,
|
||||
):
|
||||
created = create_db_engine("postgresql+asyncpg://user:pass@db/node")
|
||||
|
||||
assert created is fake_engine
|
||||
assert factory.call_args.kwargs["pool_pre_ping"] is True
|
||||
assert listen.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_created_engine_warns_for_long_checkouts(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: object
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_hold_warn_seconds", 0.0)
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
first = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/first.db")
|
||||
second = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/second.db")
|
||||
|
||||
try:
|
||||
with patch.object(db.logger, "warning") as warning:
|
||||
async with first.connect() as connection:
|
||||
await connection.exec_driver_sql("SELECT 1")
|
||||
async with second.connect() as connection:
|
||||
await connection.exec_driver_sql("SELECT 1")
|
||||
|
||||
assert warning.call_count == 2
|
||||
assert all(
|
||||
call.kwargs["extra"]["threshold_seconds"] == 0.0
|
||||
for call in warning.call_args_list
|
||||
)
|
||||
finally:
|
||||
await first.dispose()
|
||||
await second.dispose()
|
||||
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import AsyncGenerator
|
||||
from typing import Any, AsyncGenerator
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
@@ -8,7 +9,8 @@ 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.auth import get_reservation_snapshot, pay_for_request
|
||||
from routstr.core.db import ApiKey, ReservationRelease
|
||||
from routstr.upstream.ehbp import (
|
||||
finalize_ehbp_actual_cost_payment,
|
||||
finalize_ehbp_max_cost_payment,
|
||||
@@ -43,18 +45,38 @@ async def _api_key(session: AsyncSession, hashed_key: str) -> ApiKey | None:
|
||||
).one_or_none()
|
||||
|
||||
|
||||
def _fail_nth_api_key_update(
|
||||
session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
target_update: int,
|
||||
) -> None:
|
||||
"""Return rowcount=0 for one API-key UPDATE without mutating the database."""
|
||||
original_exec = session.exec
|
||||
api_key_updates = 0
|
||||
|
||||
async def exec_with_failure(
|
||||
statement: Any, *args: Any, **kwargs: Any
|
||||
) -> Any:
|
||||
nonlocal api_key_updates
|
||||
table = getattr(statement, "table", None)
|
||||
if getattr(table, "name", None) == "api_keys":
|
||||
api_key_updates += 1
|
||||
if api_key_updates == target_update:
|
||||
return MagicMock(rowcount=0)
|
||||
return await original_exec(statement, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(session, "exec", exec_with_failure)
|
||||
|
||||
|
||||
@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,
|
||||
)
|
||||
key = ApiKey(hashed_key="ehbp-actual", balance=10_000)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
await pay_for_request(key, 3_000, session)
|
||||
reservation = await get_reservation_snapshot(key, session)
|
||||
|
||||
await finalize_ehbp_actual_cost_payment(
|
||||
key,
|
||||
@@ -68,6 +90,7 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve
|
||||
"input_msats": 500,
|
||||
"output_msats": 700,
|
||||
},
|
||||
reservation_snapshot=reservation,
|
||||
)
|
||||
|
||||
updated = await _api_key(session, "ehbp-actual")
|
||||
@@ -82,28 +105,22 @@ async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve
|
||||
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,
|
||||
)
|
||||
parent = ApiKey(hashed_key="ehbp-parent", balance=10_000)
|
||||
child = ApiKey(
|
||||
hashed_key="ehbp-child",
|
||||
balance=0,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
parent_key_hash="ehbp-parent",
|
||||
hashed_key="ehbp-child", balance=0, parent_key_hash="ehbp-parent"
|
||||
)
|
||||
session.add(parent)
|
||||
session.add(child)
|
||||
await session.commit()
|
||||
await pay_for_request(child, 3_000, session)
|
||||
reservation = await get_reservation_snapshot(child, session)
|
||||
|
||||
await finalize_ehbp_max_cost_payment(
|
||||
child,
|
||||
session,
|
||||
max_cost_for_model=3_000,
|
||||
model_id="tinfoil/model",
|
||||
reservation_snapshot=reservation,
|
||||
)
|
||||
|
||||
updated_parent = await _api_key(session, "ehbp-parent")
|
||||
@@ -123,17 +140,16 @@ async def test_finalize_max_cost_payment_updates_parent_and_child_spend(
|
||||
@pytest.mark.asyncio
|
||||
async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matches_no_rows(
|
||||
session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
key = ApiKey(
|
||||
hashed_key="ehbp-missing-parent",
|
||||
balance=10_000,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
)
|
||||
key = ApiKey(hashed_key="ehbp-missing-parent", balance=10_000)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
await session.delete(key)
|
||||
await session.commit()
|
||||
await pay_for_request(key, 3_000, session)
|
||||
reservation = await get_reservation_snapshot(key, session)
|
||||
_fail_nth_api_key_update(session, monkeypatch, target_update=1)
|
||||
rollback_spy = AsyncMock(wraps=session.rollback)
|
||||
monkeypatch.setattr(session, "rollback", rollback_spy)
|
||||
|
||||
await finalize_ehbp_actual_cost_payment(
|
||||
key,
|
||||
@@ -141,45 +157,52 @@ async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matche
|
||||
reserved_cost_for_model=3_000,
|
||||
model_id="tinfoil/model",
|
||||
cost_info={"total_msats": 1_200},
|
||||
reservation_snapshot=reservation,
|
||||
)
|
||||
|
||||
assert await _api_key(session, "ehbp-missing-parent") is None
|
||||
rollback_spy.assert_awaited_once()
|
||||
updated = await _api_key(session, "ehbp-missing-parent")
|
||||
assert updated is not None
|
||||
assert updated.balance == 10_000
|
||||
assert updated.reserved_balance == 3_000
|
||||
assert updated.total_spent == 0
|
||||
release = await session.get(ReservationRelease, reservation.release_id)
|
||||
assert release is not None
|
||||
assert release.status == "active"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_finalize_max_cost_payment_rolls_back_parent_when_child_update_matches_no_rows(
|
||||
session: AsyncSession,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
parent = ApiKey(
|
||||
hashed_key="ehbp-rollback-parent",
|
||||
balance=10_000,
|
||||
reserved_balance=3_000,
|
||||
reserved_at=123,
|
||||
)
|
||||
parent = ApiKey(hashed_key="ehbp-rollback-parent", balance=10_000)
|
||||
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 pay_for_request(child, 3_000, session)
|
||||
reservation = await get_reservation_snapshot(child, session)
|
||||
_fail_nth_api_key_update(session, monkeypatch, target_update=2)
|
||||
|
||||
await finalize_ehbp_max_cost_payment(
|
||||
child,
|
||||
session,
|
||||
max_cost_for_model=3_000,
|
||||
model_id="tinfoil/model",
|
||||
reservation_snapshot=reservation,
|
||||
)
|
||||
|
||||
updated_parent = await _api_key(session, "ehbp-rollback-parent")
|
||||
assert updated_parent is not None
|
||||
assert updated_parent.balance == 10_000
|
||||
assert updated_parent.reserved_balance == 3_000
|
||||
assert updated_parent.reserved_at == 123
|
||||
assert updated_parent.total_spent == 0
|
||||
assert await _api_key(session, "ehbp-missing-child") is None
|
||||
updated_child = await _api_key(session, "ehbp-missing-child")
|
||||
assert updated_child is not None
|
||||
assert updated_child.reserved_balance == 3_000
|
||||
assert updated_child.total_spent == 0
|
||||
|
||||
@@ -0,0 +1,388 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncGenerator
|
||||
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
|
||||
|
||||
|
||||
class _SessionContext:
|
||||
def __init__(self, session: Mock) -> None:
|
||||
self.session = session
|
||||
|
||||
async def __aenter__(self) -> Mock:
|
||||
return self.session
|
||||
|
||||
async def __aexit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _session_context(session: Mock) -> _SessionContext:
|
||||
return _SessionContext(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_prepares_wallet_then_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 prepare(*_args: object, **_kwargs: object) -> Mock:
|
||||
events.append("prepare")
|
||||
return payout_wallet
|
||||
|
||||
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(side_effect=prepare)),
|
||||
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 == ["prepare", "checkpoint", "send", "complete"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_preparation_failure_does_not_checkpoint() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
checkpoint = 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", checkpoint),
|
||||
patch(
|
||||
"routstr.wallet.get_wallet",
|
||||
AsyncMock(side_effect=RuntimeError("wallet unavailable")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
checkpoint.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
send = 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=False),
|
||||
),
|
||||
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", send),
|
||||
patch("routstr.wallet.logger.warning") as warning,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
send.assert_not_awaited()
|
||||
warning.assert_called_once_with("Routstr fee payout was already claimed")
|
||||
|
||||
|
||||
@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()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_cancellation_during_send_alerts_and_propagates() -> 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(return_value=None)),
|
||||
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=asyncio.CancelledError()),
|
||||
),
|
||||
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()
|
||||
assert critical.call_args.args[0] == (
|
||||
"Routstr fee payout outcome is unknown; manual reconciliation required"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure_site", ["session", "completion"])
|
||||
async def test_fee_payout_completion_failures_use_sent_checkpoint_alert(
|
||||
failure_site: str,
|
||||
) -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
completion = AsyncMock()
|
||||
if failure_site == "session":
|
||||
create_session = Mock(
|
||||
side_effect=[
|
||||
_session_context(session),
|
||||
_session_context(session),
|
||||
RuntimeError("pool unavailable"),
|
||||
]
|
||||
)
|
||||
else:
|
||||
create_session = Mock(return_value=_session_context(session))
|
||||
completion.side_effect = RuntimeError("checkpoint unavailable")
|
||||
|
||||
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", create_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", completion),
|
||||
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(return_value=5)),
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
critical.assert_called_once()
|
||||
assert critical.call_args.args[0] == (
|
||||
"Routstr fee payout sent but checkpoint was not completed"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_releases_db_connection_during_send(tmp_path: object) -> None:
|
||||
"""With pool_size=1, the payout must not hold a connection while the
|
||||
external LNURL send is in flight, or the completion step would starve."""
|
||||
engine = create_async_engine(
|
||||
f"sqlite+aiosqlite:///{tmp_path}/payout.db", pool_size=1, max_overflow=0
|
||||
)
|
||||
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_000))
|
||||
await session.commit()
|
||||
|
||||
@asynccontextmanager
|
||||
async def create_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
yield session
|
||||
|
||||
async def send(*_args: object, **_kwargs: object) -> int:
|
||||
assert engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
return 5
|
||||
|
||||
try:
|
||||
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", create_session),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
fee = await db.get_routstr_fee(session)
|
||||
assert fee.payout_in_progress_msats == 0
|
||||
assert fee.total_paid_msats == 5_000_000
|
||||
finally:
|
||||
await engine.dispose()
|
||||
@@ -0,0 +1,116 @@
|
||||
import os
|
||||
import sqlite3
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from alembic.config import Config
|
||||
from alembic.script import ScriptDirectory
|
||||
|
||||
|
||||
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_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None:
|
||||
root = Path(__file__).resolve().parents[2]
|
||||
database_path = tmp_path / "fresh-node.db"
|
||||
database_url = f"sqlite+aiosqlite:///{database_path}"
|
||||
|
||||
_run_alembic(root, database_url, "head")
|
||||
|
||||
with sqlite3.connect(database_path) as connection:
|
||||
version = connection.execute(
|
||||
"SELECT version_num FROM alembic_version"
|
||||
).fetchone()
|
||||
columns = {
|
||||
row[1] for row in connection.execute("PRAGMA table_info(routstr_fees)")
|
||||
}
|
||||
fee = connection.execute(
|
||||
"SELECT id, accumulated_msats, total_paid_msats, last_paid_at, "
|
||||
"payout_in_progress_msats, payout_started_at FROM routstr_fees"
|
||||
).fetchone()
|
||||
|
||||
migration_config = Config(str(root / "alembic.ini"))
|
||||
assert version == (
|
||||
ScriptDirectory.from_config(migration_config).get_current_head(),
|
||||
)
|
||||
assert {
|
||||
"id",
|
||||
"accumulated_msats",
|
||||
"total_paid_msats",
|
||||
"last_paid_at",
|
||||
"payout_in_progress_msats",
|
||||
"payout_started_at",
|
||||
} <= columns
|
||||
assert fee == (1, 0, 0, None, 0, None)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
def test_fee_payout_checkpoint_repair_restores_columns_missing_at_old_head(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
root = Path(__file__).resolve().parents[2]
|
||||
database_path = tmp_path / "migration.db"
|
||||
database_url = f"sqlite+aiosqlite:///{database_path}"
|
||||
old_head = "7f2843d3f4e4"
|
||||
_run_alembic(root, database_url, old_head)
|
||||
|
||||
# Reproduce a database that was stamped to head after a duplicate-column or
|
||||
# unknown-revision recovery skipped part of the migration chain.
|
||||
with sqlite3.connect(database_path) as connection:
|
||||
connection.execute("ALTER TABLE routstr_fees DROP COLUMN payout_started_at")
|
||||
connection.execute(
|
||||
"ALTER TABLE routstr_fees DROP COLUMN payout_in_progress_msats"
|
||||
)
|
||||
connection.commit()
|
||||
|
||||
_run_alembic(root, database_url, "head")
|
||||
|
||||
with sqlite3.connect(database_path) as connection:
|
||||
columns = {
|
||||
row[1] for row in connection.execute("PRAGMA table_info(routstr_fees)")
|
||||
}
|
||||
row = connection.execute(
|
||||
"SELECT payout_in_progress_msats, payout_started_at "
|
||||
"FROM routstr_fees WHERE id = 1"
|
||||
).fetchone()
|
||||
|
||||
assert {"payout_in_progress_msats", "payout_started_at"} <= columns
|
||||
assert row == (0, None)
|
||||
@@ -1,17 +1,41 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncGenerator, Generator
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.wallet import fetch_all_balances
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_balance_fetch_state() -> Generator[None, None, None]:
|
||||
from routstr import wallet
|
||||
|
||||
wallet._balance_fetch_failures.clear()
|
||||
wallet._balance_fetch_locks.clear()
|
||||
wallet._mint_supported_units.clear()
|
||||
wallet._MintRateGuard._guards.clear()
|
||||
yield
|
||||
wallet._balance_fetch_failures.clear()
|
||||
wallet._balance_fetch_locks.clear()
|
||||
wallet._mint_supported_units.clear()
|
||||
wallet._MintRateGuard._guards.clear()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
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())),
|
||||
@@ -21,11 +45,13 @@ def _patches(proof_amount: int = 1000): # type: ignore[no-untyped-def]
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
AsyncMock(side_effect=lambda proofs, wallet, **kwargs: proofs),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_for_mint_and_unit",
|
||||
AsyncMock(return_value=0),
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(
|
||||
return_value={("http://primary:3338", "sat"): user_balance_msats}
|
||||
),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
]
|
||||
@@ -36,8 +62,9 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None:
|
||||
"""With empty cashu_mints, balances are still fetched for primary_mint."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with patch.object(settings, "cashu_mints", []), patch.object(
|
||||
settings, "primary_mint", "http://primary:3338"
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
):
|
||||
for p in _patches(proof_amount=1000):
|
||||
p.start()
|
||||
@@ -52,14 +79,341 @@ 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_uses_units_advertised_by_mint() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch(
|
||||
"routstr.wallet._get_supported_mint_units",
|
||||
AsyncMock(return_value=["sat"]),
|
||||
) as supported_units,
|
||||
):
|
||||
for p in _patches(proof_amount=1000):
|
||||
p.start()
|
||||
try:
|
||||
details, *_ = await fetch_all_balances()
|
||||
finally:
|
||||
patch.stopall()
|
||||
|
||||
supported_units.assert_awaited_once_with("http://mint:3338")
|
||||
assert [detail["unit"] for detail in details] == ["sat"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unit_discovery_failure_returns_structured_balance_error() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
get_wallet = AsyncMock()
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch(
|
||||
"routstr.wallet._get_supported_mint_units",
|
||||
AsyncMock(side_effect=httpx.ConnectError("mint unavailable")),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
):
|
||||
details, *_ = await fetch_all_balances()
|
||||
|
||||
assert details[0]["unit"] == settings.primary_mint_unit
|
||||
assert details[0]["error_code"] == "unreachable"
|
||||
assert details[0]["retry_after_seconds"] > 0
|
||||
get_wallet.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_supported_mint_units_come_from_active_keysets() -> None:
|
||||
from routstr.core.settings import settings
|
||||
from routstr.wallet import _get_supported_mint_units
|
||||
|
||||
# Cashu versions/mints may deserialize keyset units as either strings or
|
||||
# Unit enum-like objects. Both representations must be accepted.
|
||||
sat = MagicMock(active=True, unit="sat")
|
||||
msat = MagicMock(active=False, unit="msat")
|
||||
usd = MagicMock(active=True)
|
||||
usd.unit.name = "usd"
|
||||
wallet = MagicMock()
|
||||
wallet._get_keysets = AsyncMock(return_value=[usd, msat, sat])
|
||||
|
||||
with (
|
||||
patch.object(settings, "primary_mint_unit", "sat"),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)),
|
||||
):
|
||||
units = await _get_supported_mint_units("http://mint:3338")
|
||||
cached_units = await _get_supported_mint_units("http://mint:3338")
|
||||
|
||||
assert units == ["sat", "usd"]
|
||||
assert cached_units == units
|
||||
wallet._get_keysets.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_backs_off_after_connection_failure() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
get_wallet = AsyncMock(side_effect=httpx.ConnectError("mint unavailable"))
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch("routstr.mint.time.monotonic", return_value=10),
|
||||
patch("routstr.wallet.logger.warning") as warning,
|
||||
):
|
||||
first = await fetch_all_balances(units=["sat"])
|
||||
second = await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert first[0][0]["error"] == "mint unavailable"
|
||||
assert first[0][0]["error_code"] == "unreachable"
|
||||
assert first[0][0]["retry_after_seconds"] == 60
|
||||
assert second[0][0]["error"] == "mint unavailable"
|
||||
assert second[0][0]["error_code"] == "unreachable"
|
||||
assert get_wallet.await_count == 1
|
||||
warning.assert_called_once()
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch("routstr.mint.time.monotonic", return_value=71),
|
||||
patch("routstr.wallet.logger.warning"),
|
||||
):
|
||||
await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert get_wallet.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_reports_rate_limit_status() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
request = httpx.Request("GET", "http://mint:3338/v1/keysets")
|
||||
response = httpx.Response(429, request=request, headers={"Retry-After": "45"})
|
||||
error = httpx.HTTPStatusError("rate limited", request=request, response=response)
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(side_effect=error)),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert details[0]["error_code"] == "rate_limited"
|
||||
assert details[0]["retry_after_seconds"] == 60
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_failure_applies_mint_cooldown_to_other_units() -> None:
|
||||
from routstr.core.settings import settings
|
||||
from routstr.wallet import _mint_cooldown_remaining
|
||||
|
||||
mint = "http://mint:3338"
|
||||
get_wallet = AsyncMock(side_effect=httpx.ConnectError("mint unavailable"))
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", [mint]),
|
||||
patch.object(settings, "primary_mint", mint),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch("routstr.mint.time.monotonic", return_value=10),
|
||||
patch("routstr.wallet.logger.warning") as warning,
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat", "msat"])
|
||||
cooldown = _mint_cooldown_remaining(mint)
|
||||
|
||||
assert get_wallet.await_count == 1
|
||||
assert warning.call_count == 1
|
||||
assert cooldown == 60
|
||||
assert details[0]["error"] == "mint unavailable"
|
||||
assert details[0]["error_code"] == "unreachable"
|
||||
assert details[1]["error"] == "Mint is unreachable"
|
||||
assert details[1]["error_code"] == "unreachable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_closes_db_session_before_concurrent_mint_io() -> None:
|
||||
"""Slow mint checks must never run while the balance DB session is open."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
session_open = False
|
||||
mint_calls = 0
|
||||
|
||||
@asynccontextmanager
|
||||
async def tracked_session(): # type: ignore[no-untyped-def]
|
||||
nonlocal session_open
|
||||
session_open = True
|
||||
try:
|
||||
yield MagicMock()
|
||||
finally:
|
||||
session_open = False
|
||||
|
||||
async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def]
|
||||
nonlocal mint_calls
|
||||
assert session_open is False
|
||||
mint_calls += 1
|
||||
await asyncio.sleep(0)
|
||||
return proofs
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://one:3338", "http://two:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://one:3338"),
|
||||
patch("routstr.wallet.db.create_session", tracked_session),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(return_value={}),
|
||||
create=True,
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=1)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=slow_filter),
|
||||
),
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat", "msat"])
|
||||
|
||||
assert mint_calls == 4
|
||||
assert all("error" not in detail for detail in details)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_bounds_parallel_mint_checks() -> None:
|
||||
"""A slow mint fleet cannot create an unbounded external-I/O fan-out."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
active = 0
|
||||
peak = 0
|
||||
|
||||
async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def]
|
||||
nonlocal active, peak
|
||||
active += 1
|
||||
peak = max(peak, active)
|
||||
await asyncio.sleep(0.01)
|
||||
active -= 1
|
||||
return proofs
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
settings,
|
||||
"cashu_mints",
|
||||
[f"http://mint-{index}:3338" for index in range(8)],
|
||||
),
|
||||
patch.object(settings, "primary_mint", ""),
|
||||
patch.object(settings, "mint_operation_concurrency", 2),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(return_value={}),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=slow_filter),
|
||||
),
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert len(details) == 8
|
||||
assert peak == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_slow_mints_do_not_exhaust_a_single_connection_pool(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Concurrent slow balance refreshes release the sole DB connection promptly."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
engine = create_async_engine(
|
||||
f"sqlite+aiosqlite:///{tmp_path / 'pool-pressure.db'}",
|
||||
pool_size=1,
|
||||
max_overflow=0,
|
||||
pool_timeout=0.2,
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
@asynccontextmanager
|
||||
async def single_pool_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
yield session
|
||||
|
||||
async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def]
|
||||
await asyncio.sleep(0.3)
|
||||
return proofs
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://slow:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://slow:3338"),
|
||||
patch.object(settings, "mint_operation_concurrency", 1),
|
||||
patch("routstr.wallet.db.create_session", single_pool_session),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=slow_filter),
|
||||
),
|
||||
):
|
||||
results = await asyncio.gather(
|
||||
*(fetch_all_balances(units=["sat"]) for _ in range(6))
|
||||
)
|
||||
|
||||
assert all("error" not in result[0][0] for result in results)
|
||||
assert engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@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."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with patch.object(
|
||||
settings, "cashu_mints", ["http://primary:3338"]
|
||||
), patch.object(settings, "primary_mint", "http://primary:3338"):
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://primary:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
):
|
||||
for p in _patches(proof_amount=1000):
|
||||
p.start()
|
||||
try:
|
||||
@@ -71,3 +425,58 @@ async def test_fetch_all_balances_no_duplicate_primary_mint() -> None:
|
||||
|
||||
assert [d["mint_url"] for d in details] == ["http://primary:3338"]
|
||||
assert total_wallet == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_degrades_when_liability_read_fails() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(side_effect=RuntimeError("db pool exhausted")),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=1000)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
),
|
||||
):
|
||||
details, total_wallet, total_user, owner = await fetch_all_balances(
|
||||
units=["sat"]
|
||||
)
|
||||
|
||||
assert details[0]["error"] == "db pool exhausted"
|
||||
assert details[0]["wallet_balance"] == 1000
|
||||
assert details[0]["user_balance"] == 0
|
||||
assert details[0]["owner_balance"] == 0
|
||||
assert (total_wallet, total_user, owner) == (1000, 0, 0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_liability_error_keeps_more_specific_mint_error() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(side_effect=RuntimeError("db pool exhausted")),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.get_wallet",
|
||||
AsyncMock(side_effect=RuntimeError("mint down")),
|
||||
),
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert details[0]["error"] == "mint down"
|
||||
|
||||
@@ -0,0 +1,370 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from cashu.core.base import MintQuoteState, Proof
|
||||
|
||||
from routstr.lightning import (
|
||||
InvoiceRecoverRequest,
|
||||
_invoice_settlement_locks,
|
||||
_is_outputs_already_signed,
|
||||
_mint_invoice_quote,
|
||||
check_invoice_payment,
|
||||
get_invoice_status,
|
||||
recover_invoice,
|
||||
)
|
||||
from routstr.wallet import Wallet
|
||||
|
||||
|
||||
def _invoice(**overrides: object) -> SimpleNamespace:
|
||||
values = {
|
||||
"id": "invoice-1",
|
||||
"payment_hash": "quote-1",
|
||||
"amount_sats": 100,
|
||||
"purpose": "create",
|
||||
"status": "pending",
|
||||
"paid_at": None,
|
||||
"api_key_hash": None,
|
||||
"mint_url": "http://mint:3338",
|
||||
"balance_limit": None,
|
||||
"balance_limit_reset": None,
|
||||
"validity_date": None,
|
||||
"created_at": 1,
|
||||
"expires_at": 2,
|
||||
}
|
||||
values.update(overrides)
|
||||
return SimpleNamespace(**values)
|
||||
|
||||
|
||||
def _proof(amount: int, mint_id: str, *, reserved: bool = False) -> Proof:
|
||||
return Proof(amount=amount, mint_id=mint_id, reserved=reserved)
|
||||
|
||||
|
||||
def _recovery_wallet(
|
||||
error: Exception,
|
||||
*,
|
||||
proofs_before: list[Proof] | None = None,
|
||||
proofs_after: list[Proof] | None = None,
|
||||
) -> Mock:
|
||||
async def load_proofs(*, reload: bool) -> None:
|
||||
if wallet.load_proofs.await_count >= 2 and proofs_after is not None:
|
||||
wallet.proofs = list(proofs_after)
|
||||
|
||||
wallet = Mock(
|
||||
mint=AsyncMock(side_effect=error),
|
||||
keysets={"keyset-1": Mock()},
|
||||
restore_tokens_for_keyset=AsyncMock(),
|
||||
load_proofs=AsyncMock(side_effect=load_proofs),
|
||||
proofs=list(proofs_before or []),
|
||||
)
|
||||
return wallet
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_mint_recovers_quote_linked_outputs_already_signed() -> None:
|
||||
invoice = _invoice()
|
||||
wallet = _recovery_wallet(
|
||||
Exception("Mint Error: outputs have already been signed before (Code: 11003)"),
|
||||
proofs_after=[_proof(100, "quote-1")],
|
||||
)
|
||||
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
wallet.restore_tokens_for_keyset.assert_awaited_once_with(
|
||||
"keyset-1", to=1, batch=25
|
||||
)
|
||||
assert wallet.load_proofs.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_mint_accepts_preloaded_quote_linked_proofs() -> None:
|
||||
invoice = _invoice()
|
||||
wallet = _recovery_wallet(
|
||||
Exception("must not mint"),
|
||||
proofs_before=[_proof(64, "quote-1"), _proof(36, "quote-1")],
|
||||
)
|
||||
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
wallet.mint.assert_not_awaited()
|
||||
wallet.restore_tokens_for_keyset.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_mint_does_not_accept_unrelated_11003_text() -> None:
|
||||
invoice = _invoice()
|
||||
error = Exception("backend request 11003 failed")
|
||||
wallet = _recovery_wallet(error)
|
||||
|
||||
with pytest.raises(Exception) as caught:
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
assert caught.value is error
|
||||
wallet.restore_tokens_for_keyset.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_installed_cashu_error_shape_recognizes_realistic_11003_phrase() -> None:
|
||||
request = httpx.Request("POST", "http://mint:3338/v1/mint/bolt11")
|
||||
response = httpx.Response(
|
||||
400,
|
||||
request=request,
|
||||
json={"detail": "outputs have already been signed before", "code": 11003},
|
||||
)
|
||||
|
||||
with pytest.raises(Exception) as caught:
|
||||
Wallet.raise_on_error_request(response)
|
||||
|
||||
assert _is_outputs_already_signed(caught.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("recovered", [0, 99])
|
||||
async def test_invoice_mint_rejects_empty_or_short_quote_recovery(
|
||||
recovered: int,
|
||||
) -> None:
|
||||
invoice = _invoice()
|
||||
wallet = _recovery_wallet(
|
||||
Exception("Mint Error: outputs already signed (Code: 11003)"),
|
||||
proofs_after=[_proof(recovered, "quote-1")] if recovered else [],
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="expected at least 100"):
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_mint_rejects_unrelated_concurrent_balance_growth() -> None:
|
||||
invoice = _invoice()
|
||||
wallet = _recovery_wallet(
|
||||
Exception("Mint Error: outputs already signed (Code: 11003)"),
|
||||
proofs_after=[_proof(10_000, "different-quote")],
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="quote-linked recovery returned 0"):
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_pending_invoice_is_not_minted() -> None:
|
||||
_invoice_settlement_locks.clear()
|
||||
invoice = _invoice(status="expired")
|
||||
session = AsyncMock()
|
||||
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock()) as get_wallet:
|
||||
await check_invoice_payment(invoice, session) # type: ignore[arg-type]
|
||||
|
||||
get_wallet.assert_not_awaited()
|
||||
session.commit.assert_awaited_once()
|
||||
assert _invoice_settlement_locks == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ambiguous_invoice_mint_timeout_remains_recoverable() -> None:
|
||||
_invoice_settlement_locks.clear()
|
||||
invoice = _invoice()
|
||||
session = AsyncMock()
|
||||
wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True)))
|
||||
state_session = AsyncMock()
|
||||
state_session.exec.return_value.rowcount = 1
|
||||
|
||||
@asynccontextmanager
|
||||
async def owned_session() -> AsyncIterator[AsyncMock]:
|
||||
yield state_session
|
||||
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.lightning.create_session", owned_session),
|
||||
patch(
|
||||
"routstr.lightning._mint_invoice_quote",
|
||||
AsyncMock(side_effect=httpx.TimeoutException("response lost")),
|
||||
),
|
||||
patch("routstr.lightning._reload_invoice_view", AsyncMock()),
|
||||
):
|
||||
await check_invoice_payment(invoice, session) # type: ignore[arg-type]
|
||||
|
||||
assert invoice.status == "settlement_pending"
|
||||
state_session.commit.assert_awaited_once()
|
||||
session.rollback.assert_not_awaited()
|
||||
# One commit closes the initial read transaction before external I/O.
|
||||
session.commit.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_quote_lookup_timeout_is_not_definitively_unpaid() -> None:
|
||||
_invoice_settlement_locks.clear()
|
||||
invoice = _invoice(expires_at=0)
|
||||
session = AsyncMock()
|
||||
wallet = Mock(
|
||||
get_mint_quote=AsyncMock(side_effect=httpx.TimeoutException("quote timeout"))
|
||||
)
|
||||
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.lightning._reload_invoice_view", AsyncMock()),
|
||||
):
|
||||
result = await check_invoice_payment(invoice, session) # type: ignore[arg-type]
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_overdue_invoice_does_not_expire_after_ambiguous_quote_lookup() -> None:
|
||||
invoice = _invoice(status="pending", expires_at=0)
|
||||
session = AsyncMock()
|
||||
session.get.return_value = invoice
|
||||
check = AsyncMock(return_value=False)
|
||||
|
||||
with patch("routstr.lightning.check_invoice_payment", check):
|
||||
response = await get_invoice_status(invoice.id, session) # type: ignore[arg-type]
|
||||
|
||||
assert response.status == "pending"
|
||||
session.commit.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_overdue_invoice_expires_only_after_definitive_unpaid_quote() -> None:
|
||||
invoice = _invoice(status="pending", expires_at=0)
|
||||
session = AsyncMock()
|
||||
session.get.return_value = invoice
|
||||
check = AsyncMock(return_value=True)
|
||||
|
||||
async def expire(
|
||||
candidate: SimpleNamespace, _session: AsyncMock, definitive: bool
|
||||
) -> bool:
|
||||
assert definitive is True
|
||||
candidate.status = "expired"
|
||||
return True
|
||||
|
||||
with (
|
||||
patch("routstr.lightning.check_invoice_payment", check),
|
||||
patch(
|
||||
"routstr.lightning._expire_invoice_if_authoritatively_unpaid",
|
||||
side_effect=expire,
|
||||
) as expire_invoice,
|
||||
):
|
||||
response = await get_invoice_status(invoice.id, session) # type: ignore[arg-type]
|
||||
|
||||
assert response.status == "expired"
|
||||
expire_invoice.assert_awaited_once_with(invoice, session, True)
|
||||
session.commit.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_recover_applies_authoritative_expiry_helper() -> None:
|
||||
invoice = _invoice(status="pending", expires_at=0)
|
||||
session = AsyncMock()
|
||||
result = Mock()
|
||||
result.first.return_value = invoice
|
||||
session.exec.return_value = result
|
||||
check = AsyncMock(return_value=True)
|
||||
|
||||
async def expire(
|
||||
candidate: SimpleNamespace, _session: AsyncMock, definitive: bool
|
||||
) -> bool:
|
||||
assert definitive is True
|
||||
candidate.status = "expired"
|
||||
return True
|
||||
|
||||
with (
|
||||
patch("routstr.lightning.check_invoice_payment", check),
|
||||
patch(
|
||||
"routstr.lightning._expire_invoice_if_authoritatively_unpaid",
|
||||
side_effect=expire,
|
||||
) as expire_invoice,
|
||||
):
|
||||
response = await recover_invoice(
|
||||
InvoiceRecoverRequest(bolt11="lnbc-test"), session # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
assert response.status == "expired"
|
||||
expire_invoice.assert_awaited_once_with(invoice, session, True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_paid_state_write_failure_still_reports_non_expirable_outcome() -> None:
|
||||
_invoice_settlement_locks.clear()
|
||||
invoice = _invoice(expires_at=0)
|
||||
session = AsyncMock()
|
||||
wallet = Mock(
|
||||
get_mint_quote=AsyncMock(
|
||||
return_value=Mock(paid=True, state=MintQuoteState.paid)
|
||||
)
|
||||
)
|
||||
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch(
|
||||
"routstr.lightning._mint_invoice_quote",
|
||||
AsyncMock(side_effect=httpx.TimeoutException("response lost")),
|
||||
),
|
||||
patch(
|
||||
"routstr.lightning.create_session",
|
||||
side_effect=RuntimeError("database unavailable"),
|
||||
),
|
||||
patch("routstr.lightning._reload_invoice_view", AsyncMock()),
|
||||
):
|
||||
definitively_unpaid = await check_invoice_payment(
|
||||
invoice, session # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
assert definitively_unpaid is False
|
||||
assert invoice.status == "pending"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_settlement_pending_invoice_does_not_expire() -> None:
|
||||
invoice = _invoice(status="settlement_pending", expires_at=0)
|
||||
session = AsyncMock()
|
||||
session.get.return_value = invoice
|
||||
check = AsyncMock()
|
||||
|
||||
with patch("routstr.lightning.check_invoice_payment", check):
|
||||
response = await get_invoice_status(
|
||||
invoice.id, session # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
check.assert_awaited_once_with(invoice, session)
|
||||
assert response.status == "settlement_pending"
|
||||
session.commit.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_invoice_checks_finalize_once_in_process() -> None:
|
||||
_invoice_settlement_locks.clear()
|
||||
invoice = _invoice()
|
||||
session = AsyncMock()
|
||||
wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True)))
|
||||
|
||||
async def refresh(obj: SimpleNamespace) -> None:
|
||||
return None
|
||||
|
||||
session.refresh = AsyncMock(side_effect=refresh)
|
||||
|
||||
@asynccontextmanager
|
||||
async def owned_session() -> AsyncIterator[AsyncMock]:
|
||||
owned = AsyncMock()
|
||||
owned.exec.return_value.rowcount = 1
|
||||
yield owned
|
||||
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.lightning.create_session", owned_session),
|
||||
patch("routstr.lightning._mint_invoice_quote", AsyncMock()),
|
||||
patch(
|
||||
"routstr.lightning._finalize_invoice_settlement",
|
||||
AsyncMock(return_value=(True, "b" * 64)),
|
||||
) as finalize,
|
||||
):
|
||||
await asyncio.gather(
|
||||
check_invoice_payment(invoice, session), # type: ignore[arg-type]
|
||||
check_invoice_payment(invoice, session), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
assert invoice.status == "paid"
|
||||
finalize.assert_awaited_once()
|
||||
assert _invoice_settlement_locks == {}
|
||||
@@ -0,0 +1,207 @@
|
||||
"""LNURL melt attempts must not misclassify ambiguous payment outcomes."""
|
||||
|
||||
import asyncio
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from cashu.core.base import MeltQuoteState
|
||||
|
||||
from routstr.core.settings import settings
|
||||
from routstr.mint import MintCooldownError, MintRateGuard
|
||||
from routstr.payment.lnurl import (
|
||||
MeltOutcomeAmbiguousError,
|
||||
raw_send_to_lnurl,
|
||||
)
|
||||
|
||||
LNURL_DATA = {
|
||||
"callback_url": "https://ln.tld/cb",
|
||||
"min_sendable": 1_000,
|
||||
"max_sendable": 100_000_000,
|
||||
}
|
||||
|
||||
|
||||
def _wallet() -> tuple[MagicMock, list[MagicMock]]:
|
||||
proofs = [MagicMock(amount=1000)]
|
||||
wallet = MagicMock(url="https://mint.test")
|
||||
wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q"))
|
||||
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
|
||||
return wallet, proofs
|
||||
|
||||
|
||||
def _lnurl_patches() -> tuple[Any, Any]:
|
||||
return (
|
||||
patch(
|
||||
"routstr.payment.lnurl.get_lnurl_data",
|
||||
AsyncMock(return_value=LNURL_DATA),
|
||||
),
|
||||
patch(
|
||||
"routstr.payment.lnurl.get_lnurl_invoice",
|
||||
AsyncMock(return_value=("lnbc1...", {})),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
|
||||
async def _hang(**kwargs: object) -> None:
|
||||
await asyncio.sleep(5)
|
||||
|
||||
wallet.melt = AsyncMock(side_effect=_hang)
|
||||
wallet.get_melt_quote = AsyncMock(
|
||||
return_value=MagicMock(state=MeltQuoteState.unpaid)
|
||||
)
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with (
|
||||
patch.object(settings, "mint_operation_timeout_seconds", 0.05),
|
||||
patch.object(settings, "mint_retry_max_attempts", 0),
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
pytest.raises(MeltOutcomeAmbiguousError, match="outcome is ambiguous"),
|
||||
):
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
wallet.get_melt_quote.assert_awaited_once_with("q")
|
||||
wallet.set_reserved_for_melt.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_timeout_reconciled_paid_is_success() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
|
||||
async def _hang(**kwargs: object) -> None:
|
||||
await asyncio.sleep(5)
|
||||
|
||||
wallet.melt = AsyncMock(side_effect=_hang)
|
||||
wallet.get_melt_quote = AsyncMock(
|
||||
return_value=MagicMock(state=MeltQuoteState.paid)
|
||||
)
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with (
|
||||
patch.object(settings, "mint_operation_timeout_seconds", 0.05),
|
||||
patch.object(settings, "mint_retry_max_attempts", 0),
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
):
|
||||
paid = await raw_send_to_lnurl(
|
||||
wallet, proofs, "owner@ln.tld", "sat", amount=1000
|
||||
)
|
||||
|
||||
assert paid > 0
|
||||
wallet.get_melt_quote.assert_awaited_once_with("q")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending))
|
||||
wallet.get_melt_quote = AsyncMock(
|
||||
return_value=MagicMock(state=MeltQuoteState.pending)
|
||||
)
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with (
|
||||
patch.object(settings, "mint_operation_timeout_seconds", 5),
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
pytest.raises(MeltOutcomeAmbiguousError, match="outcome is ambiguous"),
|
||||
):
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
wallet.get_melt_quote.assert_awaited_once_with("q")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("rate_error", ["cooldown", "http_429"])
|
||||
async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs(
|
||||
rate_error: str,
|
||||
) -> None:
|
||||
wallet, proofs = _wallet()
|
||||
wallet.melt = AsyncMock()
|
||||
wallet.set_reserved_for_send = AsyncMock()
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
async def run_operation(factory: Any, *, op_name: str, **_: object) -> Any:
|
||||
if op_name == "lnurl_melt":
|
||||
if rate_error == "cooldown":
|
||||
raise MintCooldownError(str(wallet.url), 60)
|
||||
request = httpx.Request("POST", f"{wallet.url}/v1/melt/bolt11")
|
||||
response = httpx.Response(429, request=request)
|
||||
raise httpx.HTTPStatusError(
|
||||
"rate limited", request=request, response=response
|
||||
)
|
||||
return await factory()
|
||||
|
||||
with (
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
patch(
|
||||
"routstr.payment.lnurl.run_mint_operation",
|
||||
side_effect=run_operation,
|
||||
),
|
||||
pytest.raises((MintCooldownError, httpx.HTTPStatusError)),
|
||||
):
|
||||
await raw_send_to_lnurl(
|
||||
wallet, proofs, "owner@ln.tld", "sat", amount=1000
|
||||
)
|
||||
|
||||
wallet.melt.assert_not_awaited()
|
||||
wallet.set_reserved_for_send.assert_awaited_once_with(
|
||||
proofs, reserved=False
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_real_mint_wrapper_http_429_unreserves_proofs() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
request = httpx.Request("POST", f"{wallet.url}/v1/melt/bolt11")
|
||||
response = httpx.Response(429, request=request)
|
||||
wallet.melt = AsyncMock(
|
||||
side_effect=httpx.HTTPStatusError(
|
||||
"rate limited", request=request, response=response
|
||||
)
|
||||
)
|
||||
wallet.set_reserved_for_send = AsyncMock()
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with (
|
||||
patch.object(settings, "mint_retry_max_attempts", 0),
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
pytest.raises(httpx.HTTPStatusError),
|
||||
):
|
||||
await raw_send_to_lnurl(
|
||||
wallet, proofs, "owner@ln.tld", "sat", amount=1000
|
||||
)
|
||||
|
||||
wallet.melt.assert_awaited_once()
|
||||
wallet.set_reserved_for_send.assert_awaited_once_with(
|
||||
proofs, reserved=False
|
||||
)
|
||||
MintRateGuard._guards.pop(str(wallet.url), None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_succeeds_on_explicit_paid_response() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid))
|
||||
wallet.get_melt_quote = AsyncMock()
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with (
|
||||
patch.object(settings, "mint_operation_timeout_seconds", 5),
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
):
|
||||
paid = await raw_send_to_lnurl(
|
||||
wallet, proofs, "owner@ln.tld", "sat", amount=1000
|
||||
)
|
||||
|
||||
assert paid > 0
|
||||
wallet.melt.assert_awaited_once()
|
||||
wallet.get_melt_quote.assert_not_awaited()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user