mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge remote-tracking branch 'origin/main' into feat/lightning-error-envelope
# Conflicts: # routstr/core/main.py
This commit is contained in:
@@ -33,6 +33,8 @@ ROUTSTR_SECRET_KEY=
|
|||||||
# DATABASE_POOL_PRE_PING=false
|
# DATABASE_POOL_PRE_PING=false
|
||||||
# Warn when a checkout is held this many seconds.
|
# Warn when a checkout is held this many seconds.
|
||||||
# DATABASE_POOL_HOLD_WARN_SECONDS=10
|
# DATABASE_POOL_HOLD_WARN_SECONDS=10
|
||||||
|
# SQLite write-lock timeout, in seconds.
|
||||||
|
# DATABASE_BUSY_TIMEOUT=30
|
||||||
# SQLite serialises writes; increasing its pool can trade pool timeouts for
|
# SQLite serialises writes; increasing its pool can trade pool timeouts for
|
||||||
# "database is locked" errors rather than increasing write throughput.
|
# "database is locked" errors rather than increasing write throughput.
|
||||||
|
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
strategy:
|
strategy:
|
||||||
matrix:
|
matrix:
|
||||||
python-version: ["3.11", "3.12"]
|
python-version: ["3.11", "3.12", "3.14"]
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
@@ -25,22 +25,22 @@ jobs:
|
|||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: |
|
run: |
|
||||||
uv sync --dev
|
uv sync --python ${{ matrix.python-version }} --dev
|
||||||
|
|
||||||
- name: Run linting with ruff
|
- name: Run linting with ruff
|
||||||
run: |
|
run: |
|
||||||
uv run ruff check .
|
uv run --python ${{ matrix.python-version }} ruff check .
|
||||||
|
|
||||||
- name: Run type checking with mypy
|
- name: Run type checking with mypy
|
||||||
run: |
|
run: |
|
||||||
uv run mypy .
|
uv run --python ${{ matrix.python-version }} mypy .
|
||||||
|
|
||||||
- name: Run tests with pytest
|
- name: Run tests with pytest
|
||||||
env:
|
env:
|
||||||
UPSTREAM_BASE_URL: "http://test"
|
UPSTREAM_BASE_URL: "http://test"
|
||||||
UPSTREAM_API_KEY: "test"
|
UPSTREAM_API_KEY: "test"
|
||||||
run: |
|
run: |
|
||||||
uv run pytest --verbose --tb=short
|
uv run --python ${{ matrix.python-version }} pytest --verbose --tb=short
|
||||||
|
|
||||||
- name: Upload test results
|
- name: Upload test results
|
||||||
if: always()
|
if: always()
|
||||||
|
|||||||
@@ -11,6 +11,9 @@ dist/
|
|||||||
*.egg
|
*.egg
|
||||||
.mypy_cache/**
|
.mypy_cache/**
|
||||||
|
|
||||||
|
# MkDocs build output
|
||||||
|
site/
|
||||||
|
|
||||||
# Development
|
# Development
|
||||||
.notes
|
.notes
|
||||||
.*keys.db
|
.*keys.db
|
||||||
|
|||||||
+1
-1
@@ -1 +1 @@
|
|||||||
3.11
|
3.14
|
||||||
|
|||||||
+3
-1
@@ -1,10 +1,12 @@
|
|||||||
FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim
|
ARG PYTHON_VERSION=3.14
|
||||||
|
FROM ghcr.io/astral-sh/uv:python${PYTHON_VERSION}-bookworm-slim
|
||||||
|
|
||||||
RUN apt-get update \
|
RUN apt-get update \
|
||||||
&& apt-get install -y --no-install-recommends \
|
&& apt-get install -y --no-install-recommends \
|
||||||
git \
|
git \
|
||||||
build-essential \
|
build-essential \
|
||||||
pkg-config \
|
pkg-config \
|
||||||
|
libffi-dev \
|
||||||
libsecp256k1-dev \
|
libsecp256k1-dev \
|
||||||
autoconf \
|
autoconf \
|
||||||
automake \
|
automake \
|
||||||
|
|||||||
+4
-1
@@ -1,4 +1,6 @@
|
|||||||
# Multi-stage Dockerfile for Routstr (includes UI build)
|
# Multi-stage Dockerfile for Routstr (includes UI build)
|
||||||
|
ARG PYTHON_VERSION=3.14
|
||||||
|
|
||||||
# Stage 1: Build the UI
|
# Stage 1: Build the UI
|
||||||
FROM node:23-alpine AS ui-builder
|
FROM node:23-alpine AS ui-builder
|
||||||
WORKDIR /app/ui
|
WORKDIR /app/ui
|
||||||
@@ -16,13 +18,14 @@ ENV NEXT_TELEMETRY_DISABLED=1
|
|||||||
RUN pnpm run build
|
RUN pnpm run build
|
||||||
|
|
||||||
# Stage 2: Build the Routstr Node
|
# Stage 2: Build the Routstr Node
|
||||||
FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim AS runner
|
FROM ghcr.io/astral-sh/uv:python${PYTHON_VERSION}-bookworm-slim AS runner
|
||||||
|
|
||||||
RUN apt-get update \
|
RUN apt-get update \
|
||||||
&& apt-get install -y --no-install-recommends \
|
&& apt-get install -y --no-install-recommends \
|
||||||
git \
|
git \
|
||||||
build-essential \
|
build-essential \
|
||||||
pkg-config \
|
pkg-config \
|
||||||
|
libffi-dev \
|
||||||
libsecp256k1-dev \
|
libsecp256k1-dev \
|
||||||
autoconf \
|
autoconf \
|
||||||
automake \
|
automake \
|
||||||
|
|||||||
@@ -307,23 +307,6 @@ ANALYTICS_KEY = os.getenv("ROUTSTR_ANALYTICS_KEY")
|
|||||||
api_key = PROD_KEY if is_production() else DEV_KEY
|
api_key = PROD_KEY if is_production() else DEV_KEY
|
||||||
```
|
```
|
||||||
|
|
||||||
### Delegated Authentication
|
|
||||||
|
|
||||||
Create sub-keys with limited permissions:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
POST /v1/wallet/create/subkey
|
|
||||||
Authorization: Bearer sk-parent-key
|
|
||||||
Content-Type: application/json
|
|
||||||
|
|
||||||
{
|
|
||||||
"name": "Limited Subkey",
|
|
||||||
"balance_limit": 1000,
|
|
||||||
"allowed_models": ["gpt-3.5-turbo"],
|
|
||||||
"expires_in_hours": 24
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
## Rate Limiting
|
## Rate Limiting
|
||||||
|
|
||||||
Rate limits are applied per API key:
|
Rate limits are applied per API key:
|
||||||
@@ -388,22 +371,6 @@ Content-Type: application/json
|
|||||||
|
|
||||||
## Monitoring
|
## Monitoring
|
||||||
|
|
||||||
### Usage Alerts
|
|
||||||
|
|
||||||
Set up usage notifications:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
POST /v1/wallet/alerts
|
|
||||||
Authorization: Bearer sk-...
|
|
||||||
Content-Type: application/json
|
|
||||||
|
|
||||||
{
|
|
||||||
"low_balance_threshold": 1000,
|
|
||||||
"daily_spend_limit": 5000,
|
|
||||||
"webhook_url": "https://your-app.com/webhook"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### Audit Logging
|
### Audit Logging
|
||||||
|
|
||||||
All API key usage is logged:
|
All API key usage is logged:
|
||||||
|
|||||||
+64
-57
@@ -418,7 +418,7 @@ POST /v1/wallet/create
|
|||||||
|
|
||||||
### Get Key Information
|
### Get Key Information
|
||||||
|
|
||||||
Get current balance, consumption data, and child keys for an API key.
|
Get current balance and consumption data for an API key.
|
||||||
|
|
||||||
```http
|
```http
|
||||||
GET /v1/balance/info
|
GET /v1/balance/info
|
||||||
@@ -432,23 +432,9 @@ Authorization: Bearer sk-...
|
|||||||
"api_key": "sk-abc...",
|
"api_key": "sk-abc...",
|
||||||
"balance": 8500000,
|
"balance": 8500000,
|
||||||
"reserved": 0,
|
"reserved": 0,
|
||||||
"is_child": false,
|
|
||||||
"parent_key": null,
|
|
||||||
"total_requests": 42,
|
"total_requests": 42,
|
||||||
"total_spent": 1500000,
|
"total_spent": 1500000,
|
||||||
"balance_limit": null,
|
"validity_date": null
|
||||||
"balance_limit_reset": null,
|
|
||||||
"validity_date": null,
|
|
||||||
"child_keys": [
|
|
||||||
{
|
|
||||||
"api_key": "sk-child1...",
|
|
||||||
"total_requests": 10,
|
|
||||||
"total_spent": 500000,
|
|
||||||
"balance_limit": 1000000,
|
|
||||||
"balance_limit_reset": "daily",
|
|
||||||
"validity_date": 1738000000
|
|
||||||
}
|
|
||||||
]
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -500,48 +486,23 @@ Authorization: Bearer sk-...
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
### Withdraw Funds
|
### Refund Balance
|
||||||
|
|
||||||
Withdraw balance as eCash.
|
Pay out the current balance. The key remains valid at zero balance and can be topped up again. The payout goes to a Lightning address when one is given (in the request or stored on the key), otherwise a Cashu token is returned.
|
||||||
|
|
||||||
```http
|
```http
|
||||||
POST /v1/wallet/withdraw
|
POST /v1/balance/refund
|
||||||
Authorization: Bearer sk-...
|
Authorization: Bearer sk-...
|
||||||
|
Content-Type: application/json
|
||||||
```
|
```
|
||||||
|
|
||||||
**Request Body:**
|
`/v1/wallet/refund` is a deprecated alias.
|
||||||
|
|
||||||
|
**Request Body** (optional):
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"amount": 5000,
|
"lightning_address": "user@getalby.com"
|
||||||
"mint": "https://mint.example.com"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**Response:**
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"cashu_token": "cashuAeyJ0...",
|
|
||||||
"amount": 5000,
|
|
||||||
"mint": "https://mint.example.com"
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
### Create Child Key
|
|
||||||
|
|
||||||
Creates one or more child API keys that share the parent's balance. Each child key creation costs a fixed amount (configurable).
|
|
||||||
|
|
||||||
```http
|
|
||||||
POST /v1/balance/child-key
|
|
||||||
Authorization: Bearer sk-...
|
|
||||||
```
|
|
||||||
|
|
||||||
**Request Body:**
|
|
||||||
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"count": 1
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -549,21 +510,67 @@ Authorization: Bearer sk-...
|
|||||||
|
|
||||||
| Parameter | Type | Required | Default | Description |
|
| Parameter | Type | Required | Default | Description |
|
||||||
|-----------|------|----------|---------|-------------|
|
|-----------|------|----------|---------|-------------|
|
||||||
| `count` | integer | Yes | - | Number of child keys to create (1-50) |
|
| `lightning_address` | string | No | Key's stored refund address | Lightning address or LNURL to pay. Overrides the stored address for this request. The effective address (request or stored) is resolved only for a request that can open a new claim, before any balance is debited. |
|
||||||
|
|
||||||
**Response:**
|
**Response (Lightning):**
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
"api_keys": ["sk-abc...", "sk-def..."],
|
"refund_id": "3f9c1e2d8b7a4c6e9f0a1b2c3d4e5f60",
|
||||||
"count": 2,
|
"status": "paid",
|
||||||
"cost_msats": 2000,
|
"recipient": "user@getalby.com",
|
||||||
"cost_sats": 2,
|
"sats": "4500"
|
||||||
"parent_balance": 98000,
|
|
||||||
"parent_balance_sats": 98
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**Response (Cashu):**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"refund_id": "3f9c1e2d8b7a4c6e9f0a1b2c3d4e5f60",
|
||||||
|
"status": "paid",
|
||||||
|
"token": "cashuAeyJ0...",
|
||||||
|
"sats": "4500"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
The amount field is `sats` or `msats` depending on the key's refund currency. It reports the gross balance debited by the claim. For Lightning refunds, mint and input fees can reduce the amount actually delivered to the recipient.
|
||||||
|
|
||||||
|
**Behaviour:**
|
||||||
|
|
||||||
|
- The balance is debited and a refund claim is recorded before the payout is attempted. A key has at most one open claim at a time.
|
||||||
|
- If the payout fails cleanly, the claim is closed and the balance is restored. Retry the request.
|
||||||
|
- Once a melt quote has been recorded or a Cashu token has been issued, the payout may already have happened, so any later failure returns `502` and withholds the balance rather than restoring it. The exception is the mint answering the melt itself with `unpaid`: that is proof nothing was sent, so the balance is restored at once and the request returns `503`.
|
||||||
|
- If the Lightning payment is dispatched but the mint cannot confirm the outcome, the request returns `502`, the balance stays withheld, and a background reconciler asks the mint until it answers. The balance is restored if the mint reports the payment unpaid.
|
||||||
|
- An unresolved claim is reported before any replay: a request on a key with an open claim returns `409` with that claim's `refund_id` and `status`.
|
||||||
|
- Calling again on a zero-balance key with no open claim returns the last paid Lightning refund, or the Cashu token issued by the last paid claim while it remains uncollected.
|
||||||
|
|
||||||
|
**Errors:**
|
||||||
|
|
||||||
|
| Status | Meaning |
|
||||||
|
|--------|---------|
|
||||||
|
| `400` | Invalid Lightning destination, no balance, or balance too small for the refund unit |
|
||||||
|
| `400` | Ongoing requests are still reserving balance on this key |
|
||||||
|
| `401` | Unknown key |
|
||||||
|
| `409` | Balance changed concurrently. Retry. |
|
||||||
|
| `409` | `refund_in_progress`: another refund claim for this key is still open. The body carries its `refund_id` and `status` |
|
||||||
|
| `409` | `refund_unresolved`: a claim for this key is `stuck` and needs operator reconciliation |
|
||||||
|
| `410` | Previously issued Cashu refund token has been swept |
|
||||||
|
| `500` | Payout failed before anything was dispatched. Balance restored. Retry. |
|
||||||
|
| `502` | Payment dispatched, outcome unconfirmed. Balance withheld pending reconciliation. Do not retry. |
|
||||||
|
| `503` | Mint unavailable, or the mint reported the Lightning payment unpaid. Balance restored. Retry later. |
|
||||||
|
|
||||||
|
**X-Cashu refunds:**
|
||||||
|
|
||||||
|
Requests paid per-call with an `X-Cashu` header get their change from this endpoint by sending the same header instead of `Authorization`:
|
||||||
|
|
||||||
|
```http
|
||||||
|
POST /v1/balance/refund
|
||||||
|
X-Cashu: cashuAeyJ0...
|
||||||
|
```
|
||||||
|
|
||||||
|
Returns the change token in the body and in an `X-Cashu` response header. `404` if no matching request exists, `425` while the change is still being minted, `410` if it was swept.
|
||||||
|
|
||||||
## Provider Discovery
|
## Provider Discovery
|
||||||
|
|
||||||
## Admin Settings
|
## Admin Settings
|
||||||
|
|||||||
+12
-8
@@ -153,28 +153,32 @@ granularity) on any of them.
|
|||||||
|--------|--------|--------|-----------|---------|
|
|--------|--------|--------|-----------|---------|
|
||||||
| `token_already_spent` | 400 | `cashu_token_already_spent` | No | The token was already redeemed. |
|
| `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. |
|
| `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_token_swap_fees_exceed_amount` | No | Token value is too small to cover the mint's NUT-02 input fees. |
|
||||||
| `mint_error` | 422 | `cashu_foreign_mint_swap_failed` | No | Swapping the token from a foreign mint to the primary mint failed. |
|
| `untrusted_mint` | 400 | `cashu_untrusted_source_mint` | No | The token was issued by a mint this node does not accept. Only the node's configured mints (`PRIMARY_MINT_URL` / `CASHU_MINTS`) are redeemable. |
|
||||||
| `mint_unreachable` | 503 | `cashu_source_mint_unreachable` | **Yes** | The mint that issued the token could not be reached; it cannot be redeemed at another mint. |
|
| `mint_unreachable` | 503 | `cashu_source_mint_unreachable` | **Yes** | The mint that issued the token could not be reached; it cannot be redeemed at another mint. |
|
||||||
| `mint_rate_limited` | 503 | `cashu_mint_rate_limited` | **Yes** | The mint rate-limited the request; retry after the cooldown. |
|
| `mint_rate_limited` | 503 | `cashu_mint_rate_limited` | **Yes** | The mint rate-limited the request; retry after the cooldown. |
|
||||||
| `mint_unreachable` | 503 | `cashu_mint_unreachable` | **Yes** | The mint could not be reached (DNS failure, refused/reset connection, timeout). The token is fine — retry once the mint recovers. |
|
| `mint_timeout` | 503 | `cashu_mint_timeout` | **Yes** | The mint did not respond in time; retry later. |
|
||||||
|
| `mint_unreachable` | 503 | `cashu_mint_unreachable` | **Yes** | The mint could not be reached (DNS failure, refused/reset connection). The token is fine — retry once the mint recovers. |
|
||||||
| `cashu_error` | 400 | `cashu_token_redemption_failed` | No | The token could not be redeemed for another expected reason. |
|
| `cashu_error` | 400 | `cashu_token_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). |
|
| `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. |
|
| `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. |
|
| `api_error` | 500 | `internal_error` | Maybe | Unexpected server-side fault during redemption. |
|
||||||
|
|
||||||
!!! important "Retry only transient mint failures"
|
!!! important "Retry only transient mint failures"
|
||||||
Only `mint_unreachable` and `mint_rate_limited` (503) are retryable — the
|
Only `mint_unreachable`, `mint_rate_limited` and `mint_timeout` (503) are
|
||||||
same token may work again later. Everything else is a permanent property of
|
retryable — the same token may work again later. Everything else is a
|
||||||
the token and must not be blindly retried. Use exponential backoff for the
|
permanent property of the token and must not be blindly retried.
|
||||||
|
`untrusted_mint` is permanent: the node will never accept that mint until
|
||||||
|
an operator adds it to `CASHU_MINTS`. Use exponential backoff for the
|
||||||
503 responses, and honor the mint's cooldown for `mint_rate_limited`. In
|
503 responses, and honor the mint's cooldown for `mint_rate_limited`. In
|
||||||
particular, a `token_consumed` 500 means the mint already spent the token,
|
particular, a `token_consumed` 500 means the mint already spent the token,
|
||||||
so a retry would fail as `token_already_spent`.
|
so a retry would fail as `token_already_spent`.
|
||||||
|
|
||||||
#### Mint failures (retryable)
|
#### Mint failures (retryable)
|
||||||
|
|
||||||
`mint_unreachable` and `mint_rate_limited` are retryable redemption errors. For
|
`mint_unreachable`, `mint_rate_limited` and `mint_timeout` are retryable
|
||||||
`mint_rate_limited`, honor the mint's cooldown before retrying.
|
redemption errors. For `mint_rate_limited`, honor the mint's cooldown before
|
||||||
|
retrying.
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -40,8 +40,9 @@ Every time you make a request to `/v1/chat/completions` (or others), the cost is
|
|||||||
`Cost = (Input_Tokens * Price_Input) + (Output_Tokens * Price_Output) + Request_Fee`
|
`Cost = (Input_Tokens * Price_Input) + (Output_Tokens * Price_Output) + Request_Fee`
|
||||||
|
|
||||||
- Prices are defined per model (see `/v1/models`).
|
- Prices are defined per model (see `/v1/models`).
|
||||||
- If you stream the response, the balance is deducted incrementally or finalized at the end of the stream.
|
- Routstr reserves an authorization ceiling before forwarding, then finalizes the request at measured token cost.
|
||||||
- If your balance hits 0 mid-stream, the connection is closed.
|
- If a successful upstream omits usage, Routstr estimates input tokens from the provider-bound request and output tokens from the returned body or streamed deltas, then applies normal model pricing.
|
||||||
|
- A reservation is only a temporary hold. Missing usage or unusable prices must never turn the full reservation into the charge; if no auditable estimate can be priced, the reservation is released without charge.
|
||||||
|
|
||||||
### Headers
|
### Headers
|
||||||
|
|
||||||
|
|||||||
@@ -300,7 +300,7 @@ Project metadata and dependencies:
|
|||||||
name = "routstr"
|
name = "routstr"
|
||||||
version = "0.2.2"
|
version = "0.2.2"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"fastapi[standard]>=0.115",
|
"fastapi[standard-no-fastapi-cloud-cli]>=0.141",
|
||||||
"sqlmodel>=0.0.24",
|
"sqlmodel>=0.0.24",
|
||||||
"cashu",
|
"cashu",
|
||||||
# ...
|
# ...
|
||||||
|
|||||||
@@ -58,10 +58,10 @@ Contains the shared opaque EHBP transport and billing helpers:
|
|||||||
- `EHBPForwardingTarget` — provider-specific target URL plus extra headers
|
- `EHBPForwardingTarget` — provider-specific target URL plus extra headers
|
||||||
- `forward_ehbp_request()` — forwards the encrypted body, captures Tinfoil
|
- `forward_ehbp_request()` — forwards the encrypted body, captures Tinfoil
|
||||||
usage from a response header or streaming HTTP trailer, and finalizes bearer
|
usage from a response header or streaming HTTP trailer, and finalizes bearer
|
||||||
billing at actual cost (falling back to max cost when usage is unavailable)
|
billing at actual cost (releasing the reservation when usage is unavailable)
|
||||||
- `forward_ehbp_x_cashu_request()` — redeems the Cashu token, refunds the full
|
- `forward_ehbp_x_cashu_request()` — redeems the Cashu token, refunds the full
|
||||||
token on upstream failure, and refunds the difference between the redeemed
|
token on upstream failure, and refunds the difference between the redeemed
|
||||||
amount and actual cost (or max cost when usage is unavailable)
|
amount and actual cost (or the full amount when usage is unavailable)
|
||||||
|
|
||||||
### Provider support
|
### Provider support
|
||||||
|
|
||||||
@@ -84,7 +84,7 @@ The proxy is a **blind relay** for EHBP requests. It cannot decrypt the body
|
|||||||
Cost tracking happens at the proxy level. Routstr reserves or redeems up to
|
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
|
`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
|
to finalize at actual token cost. If trusted usage is missing or invalid, the
|
||||||
proxy safely falls back to max-cost billing.
|
proxy releases/refunds rather than treating the authorization ceiling as usage.
|
||||||
|
|
||||||
## End-to-end flow
|
## End-to-end flow
|
||||||
|
|
||||||
|
|||||||
@@ -26,6 +26,20 @@ If you want to run a node, resell API access, or monetize hardware.
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## 👥 For Teams (Remote Nodes)
|
||||||
|
|
||||||
|
If you want one shared Routstr endpoint for a whole team, with per-member identities and per-member usage tracking.
|
||||||
|
|
||||||
|
- **[Overview](teams/index.md)**: What a remote node is, and when to use one.
|
||||||
|
- **[Deploy on Cloudron](teams/deploy-cloudron.md)**: The packaged, supported deployment.
|
||||||
|
- **[Deploy with Docker](teams/deploy-docker.md)**: Run it on any host behind your own TLS.
|
||||||
|
- **[Team Members](teams/team-members.md)**: Bootstrap the first admin and invite people.
|
||||||
|
- **[Connecting Clients](teams/clients.md)**: Wire up Claude Code, Pi, OpenCode, and API keys.
|
||||||
|
- **[Usage and Model Policy](teams/usage-and-policy.md)**: Per-member spend and model allowlists.
|
||||||
|
- **[Security Model](teams/security.md)**: Auth rules and endpoint scoping.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## 🔌 API Reference
|
## 🔌 API Reference
|
||||||
|
|
||||||
- **[Overview](api/overview.md)**: Base URL, headers, and standards.
|
- **[Overview](api/overview.md)**: Base URL, headers, and standards.
|
||||||
|
|||||||
@@ -9,6 +9,19 @@ Routstr uses **Nostr** as a decentralized directory for service discovery. Your
|
|||||||
1. **Provider Advertisement (Kind 38421)**: Your node periodically publishes an event with its URL, models, and pricing
|
1. **Provider Advertisement (Kind 38421)**: Your node periodically publishes an event with its URL, models, and pricing
|
||||||
2. **Client Discovery**: Clients query relays for these events to find suitable providers
|
2. **Client Discovery**: Clients query relays for these events to find suitable providers
|
||||||
|
|
||||||
|
### When announcements are published
|
||||||
|
|
||||||
|
Your node publishes an advertisement as soon as it has both a **Nsec** and at least one
|
||||||
|
reachable endpoint (a public `HTTP_URL`, or an `.onion` address). Saving the Nsec in the
|
||||||
|
dashboard is enough — the announcement follows within a minute, and **no restart is
|
||||||
|
required**. After the first publish it re-announces every 24 hours, and immediately
|
||||||
|
whenever the Nsec, endpoints, mints or relays change.
|
||||||
|
|
||||||
|
A node that has no Nsec yet simply waits, and starts announcing the moment one is
|
||||||
|
configured. Note that `HTTP_URL` defaults to `http://localhost:8000`, which is not a
|
||||||
|
reachable endpoint: a node with the default value and no onion address has nothing to
|
||||||
|
advertise and will not publish until one is set.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Configuration
|
## Configuration
|
||||||
|
|||||||
@@ -0,0 +1,153 @@
|
|||||||
|
# Connecting Clients
|
||||||
|
|
||||||
|
A **client** is one agent or application talking to the node. Each client gets its own ID and its own API key (`sk-...`). Clients are how a team node attributes usage: every request is billed against the client that made it.
|
||||||
|
|
||||||
|
Members create and manage **their own** clients. They cannot see or touch anyone else's.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Two credentials, two jobs
|
||||||
|
|
||||||
|
The node accepts two entirely different kinds of credential, and confusing them is the most common source of `403`s.
|
||||||
|
|
||||||
|
| Credential | Header | Purpose | Can do |
|
||||||
|
|---|---|---|---|
|
||||||
|
| **API key** | `Authorization: Bearer sk-...` | Inference | Send chat/completion requests. Cannot touch wallets, clients, or npubs. |
|
||||||
|
| **NIP-98** | `Authorization: Nostr <base64-event>` | Management | Manage clients, npubs, wallet, node control. Signed per-request by the member's `nsec`. |
|
||||||
|
|
||||||
|
Your agents use the **API key**. The `routstrd` CLI uses **NIP-98** automatically, which is why it needs your `nsec` in `~/.routstrd/config.json`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Add a client
|
||||||
|
|
||||||
|
From the member's own machine:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Name it explicitly
|
||||||
|
routstrd clients add --name "My Laptop"
|
||||||
|
|
||||||
|
# Or use a one-shot integration setup
|
||||||
|
routstrd clients add --claude-code
|
||||||
|
routstrd clients add --pi-agent
|
||||||
|
routstrd clients add --opencode
|
||||||
|
routstrd clients add --openclaw
|
||||||
|
routstrd clients add --hermes
|
||||||
|
```
|
||||||
|
|
||||||
|
The integration flags configure the agent's own config file as well as registering the client, so you do not have to hand-edit anything. Several can be combined in one call.
|
||||||
|
|
||||||
|
On success the CLI prints the credentials and the endpoint to point at:
|
||||||
|
|
||||||
|
```text
|
||||||
|
Client created.
|
||||||
|
|
||||||
|
ID: my-laptop
|
||||||
|
Name: My Laptop
|
||||||
|
API Key: sk-9f2a...
|
||||||
|
|
||||||
|
Access Routstr at: https://team.example.com/v1
|
||||||
|
```
|
||||||
|
|
||||||
|
!!! warning "The API key is a secret"
|
||||||
|
Treat `sk-...` like a password. It bills inference to the team wallet. Do not commit it, and do not paste it into a chat — unlike an npub, it is not safe to share.
|
||||||
|
|
||||||
|
### Adding is idempotent
|
||||||
|
|
||||||
|
Running `clients add` with a name that already exists does not create a duplicate. It looks the client up and prints the existing record — including its API key — so re-running is a safe way to recover a key you lost:
|
||||||
|
|
||||||
|
```text
|
||||||
|
Client 'my-laptop' already exists.
|
||||||
|
|
||||||
|
ID: my-laptop
|
||||||
|
Name: My Laptop
|
||||||
|
API Key: sk-9f2a...
|
||||||
|
```
|
||||||
|
|
||||||
|
### List and delete
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd clients list
|
||||||
|
routstrd clients delete my-laptop
|
||||||
|
```
|
||||||
|
|
||||||
|
`clients list` shows **only your own** clients.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Going beyond the CLI
|
||||||
|
|
||||||
|
The agent integrations cover the common tools, but any OpenAI-compatible client works — point it at the node and use the API key:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl https://team.example.com/v1/chat/completions \
|
||||||
|
-H "Authorization: Bearer sk-9f2a..." \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d '{
|
||||||
|
"model": "gpt-4o-mini",
|
||||||
|
"messages": [{"role": "user", "content": "hello"}]
|
||||||
|
}'
|
||||||
|
```
|
||||||
|
|
||||||
|
The base URL is always the node host plus `/v1`. Discover available models without any credential at all:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl https://team.example.com/v1/models
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## How client ownership works
|
||||||
|
|
||||||
|
This is the part that explains the odd-looking IDs on the node.
|
||||||
|
|
||||||
|
### IDs are derived from the name
|
||||||
|
|
||||||
|
A client's ID is the name lowercased with internal whitespace collapsed to hyphens, then stripped of anything that is not alphanumeric or a hyphen. `"My Laptop!"` becomes `my-laptop`.
|
||||||
|
|
||||||
|
### The node appends an owner suffix
|
||||||
|
|
||||||
|
So that two members can both have a client called `my-laptop` without colliding, the auth proxy appends the **last 7 characters of the owner's npub** to the ID before it reaches the daemon:
|
||||||
|
|
||||||
|
| Where you look | Client ID |
|
||||||
|
|---|---|
|
||||||
|
| Member's `routstrd clients list` | `my-laptop` |
|
||||||
|
| On the node (`cloudron exec`, then `routstrd clients list`) | `my-laptop-4f2x9k7` |
|
||||||
|
|
||||||
|
The suffix is stripped again on the way back, so members always see the clean ID. On the node you deliberately see the suffixed form — **the trailing characters are what tell you which member owns a client.**
|
||||||
|
|
||||||
|
### Ownership is recorded explicitly
|
||||||
|
|
||||||
|
Newly created clients store the owner's npub in an `ownerNpub` field, and the proxy authorises against that field. Clients created before that field existed fall back to matching the ID suffix, so older installs keep working until those clients are recreated.
|
||||||
|
|
||||||
|
### Consequence: admins are not automatically superusers here
|
||||||
|
|
||||||
|
`/clients`, `/clients/add`, and `/clients/delete` are owner-scoped by the calling npub. Even an `admin` cannot list or delete a colleague's clients through these endpoints. Cross-member visibility comes from running the CLI **on the node itself**, where the daemon is unauthenticated on loopback:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cloudron exec --app routstr.example.com
|
||||||
|
routstrd clients list # all clients, all owners, suffixed IDs
|
||||||
|
```
|
||||||
|
|
||||||
|
That is also the only practical way to clean up a departing member's keys — see [Team Members](team-members.md#what-revocation-does-and-does-not-do).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Refreshing models and integrations
|
||||||
|
|
||||||
|
The `clients` command carries options for the daemon's scheduled refresh job, which updates the Routstr 21 model list and re-syncs client integrations:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd clients --manual-refresh # refresh now, once
|
||||||
|
routstrd clients --disable-automatic-refresh # stop the scheduled job
|
||||||
|
routstrd clients --enable-automatic-refresh # start it again
|
||||||
|
```
|
||||||
|
|
||||||
|
The model list matters because it is also what the [model allowlist](usage-and-policy.md#model-allowlist) is enforced against.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Next steps
|
||||||
|
|
||||||
|
- [Usage and Model Policy](usage-and-policy.md) — watch what those clients are spending.
|
||||||
|
- [Security Model](security.md) — the exact rules applied to each credential.
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
# Deploy on Cloudron
|
||||||
|
|
||||||
|
[Cloudron](https://www.cloudron.io/) is the supported deployment target for a team node. The packaged image already contains **both** processes — the `routstrd` daemon and the `routstrd-auth` proxy — supervised inside a single container, with `/app/data` handled as persistent storage and TLS terminated by the platform.
|
||||||
|
|
||||||
|
| | |
|
||||||
|
|---|---|
|
||||||
|
| **App ID** | `io.routstr.routstrd-auth` |
|
||||||
|
| **Public port** | `8008` (Cloudron proxies it over HTTPS on 443) |
|
||||||
|
| **Health check** | `GET /health` |
|
||||||
|
| **Memory limit** | 512 MB |
|
||||||
|
| **Minimum box version** | Cloudron 9.1.0 |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Prerequisites
|
||||||
|
|
||||||
|
- A running Cloudron box with a domain that can get a certificate.
|
||||||
|
- The [`cloudron` CLI](https://docs.cloudron.io/cli/) installed and logged in, **only if** you are building the image yourself:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
npm install -g cloudron
|
||||||
|
cloudron login my.example.com
|
||||||
|
```
|
||||||
|
|
||||||
|
## Install
|
||||||
|
|
||||||
|
### Option A — from the published version list
|
||||||
|
|
||||||
|
The app is published as a custom Cloudron app with a version list (`CloudronVersions.json`), currently at `0.1.26`. Once that app store entry is registered on your Cloudron instance, install it from the dashboard, or:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cloudron install --appstore-id io.routstr.routstrd-auth --location routstr.example.com
|
||||||
|
```
|
||||||
|
|
||||||
|
### Option B — build the image yourself
|
||||||
|
|
||||||
|
Use this when you want to run a local modification:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git clone https://github.com/routstr/routstrd-remote
|
||||||
|
cd routstrd-remote
|
||||||
|
|
||||||
|
cloudron build # builds the Dockerfile and pushes it to your registry
|
||||||
|
cloudron install --image <registry>/routstrd-remote:<tag> --location routstr.example.com
|
||||||
|
```
|
||||||
|
|
||||||
|
!!! note "The Dockerfile is the Cloudron image"
|
||||||
|
The repository's `Dockerfile` is built `FROM cloudron/base:5.0.0` and its `CMD` is `cloudron/start.sh`, which prepares `/app/data` and starts `supervisord`. It expects Cloudron's filesystem conventions and should not be confused with a generic Docker image. See [Deploy with Docker](deploy-docker.md) for what that means in practice.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Bootstrap the first admin
|
||||||
|
|
||||||
|
The moment the app is healthy, the npub table is **empty**, and nothing except the public endpoints can be reached. Claim it before anything else — while the table is empty, `POST /npubs` is accepted without authentication, so this is the only window in which an unauthenticated registration succeeds.
|
||||||
|
|
||||||
|
On the machine of whoever will be the first admin:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bun i -g routstrd
|
||||||
|
routstrd remote https://routstr.example.com
|
||||||
|
routstrd npubs register --name "Alice"
|
||||||
|
```
|
||||||
|
|
||||||
|
`routstrd remote` generates a fresh Nostr identity if you do not have one, stores it in `~/.routstrd/config.json`, and prints your npub. `routstrd npubs register` then posts that npub and, because no npubs exist yet, receives `admin`.
|
||||||
|
|
||||||
|
!!! warning "Register immediately after install"
|
||||||
|
Until the first admin registers, anyone who knows the URL can claim the node. Do this as part of the install, not later.
|
||||||
|
|
||||||
|
Verify:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd npubs list
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
Cloudron defaults are set by `cloudron/start.sh` and the two supervisor programs. Everything below can be overridden through the Cloudron **Environment Variables** tab.
|
||||||
|
|
||||||
|
| Variable | Default | Purpose |
|
||||||
|
|---|---|---|
|
||||||
|
| `ROUTSTRD_AUTH_PORT` | `8008` | Public port served by the auth proxy. Must match the manifest's `httpPort`. |
|
||||||
|
| `ROUTSTRD_AUTH_HOST` | `0.0.0.0` | Bind address of the auth proxy. |
|
||||||
|
| `ROUTSTRD_UPSTREAM` | `http://localhost:8009` | Where the daemon listens. Keep this on loopback. |
|
||||||
|
| `ROUTSTRD_PORT` | `8009` | Port the daemon binds. |
|
||||||
|
| `ROUTSTRD_DIR` | `/app/data/routstrd` | Config directory shared by the daemon and the proxy. |
|
||||||
|
| `ROUTSTRD_DB_PATH` | `/app/data/routstrd/routstr.db` | Shared SQLite database. |
|
||||||
|
| `ROUTSTRD_CONFIG_FILE` | `$ROUTSTRD_DIR/config.json` | Daemon config file. |
|
||||||
|
| `ROUTSTRD_AUTH_MODEL_ALLOWLIST` | `false` | Set to `true` to restrict the team to the Routstr 21 model list. See [Usage and Model Policy](usage-and-policy.md). |
|
||||||
|
| `ROUTSTRD_AUTH_ADMIN_NPUBS` | *(unset)* | Optional bootstrap admins. See below. |
|
||||||
|
|
||||||
|
### Bootstrapping admins from the environment
|
||||||
|
|
||||||
|
Instead of the interactive `npubs register` step you can seed admins declaratively. Three variables are accepted and merged: `ROUTSTRD_AUTH_ADMIN_NPUBS`, `ROUTSTRD_AUTH_ADMIN_PUBKEYS`, and `ROUTSTRD_AUTH_BOOTSTRAP_NPUB`. Values are comma- or whitespace-separated and may be either `npub1...` or 64-character hex.
|
||||||
|
|
||||||
|
Rows created this way are tagged `source = 'env'`. At every startup the proxy **reconciles** them: an env-sourced row whose pubkey is no longer present in the environment is **deleted**. This means the environment variables are the source of truth for those rows — removing someone from the variable revokes their access on the next restart.
|
||||||
|
|
||||||
|
!!! tip "Prefer `npubs register` for the first admin"
|
||||||
|
There is deliberately **no** hardcoded default admin npub in the image. An image with a baked-in admin pubkey would hand control of every deployment to the same key.
|
||||||
|
|
||||||
|
### Filesystem layout
|
||||||
|
|
||||||
|
| Path | Lifetime | Contents |
|
||||||
|
|---|---|---|
|
||||||
|
| `/app/code` | replaced on update | Auth proxy source and the `start.sh` / `run-auth.sh` scripts. |
|
||||||
|
| `/app/data` | persistent, backed up | `routstrd/config.json`, `routstrd/routstr.db`, `logs/`, and a `.initialized` marker. |
|
||||||
|
| `/run` | ephemeral | `supervisord` socket and pid file. |
|
||||||
|
|
||||||
|
The startup script writes `authUrl` into `routstrd/config.json` pointing at the local proxy, and generates the container's Nostr identity (`nsec`) on first boot if one is missing. That generated identity is what authorises the daemon's own calls back through the proxy.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Backups
|
||||||
|
|
||||||
|
Cloudron's `localstorage` addon makes `/app/data` persistent and includes it in regular backups. The entire data directory is covered, which means the SQLite database at `/app/data/routstrd/routstr.db` travels with it.
|
||||||
|
|
||||||
|
To restore, restore the app from a Cloudron backup — the wallet configuration, npub table, and client records all come back together. Do not hand-copy files between hosts: the daemon's `config.json` holds the container's `nsec`, and losing it breaks the node's ability to authenticate to its own proxy.
|
||||||
|
|
||||||
|
!!! warning "Use `cloudron exec` carefully"
|
||||||
|
The database uses SQLite WAL mode. If you need a manual snapshot, stop the app first (`cloudron stop`) or use `sqlite3 .backup` inside the container — copying the `.db` file while the daemon is writing can produce an inconsistent file.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Updates
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cloudron update --app routstr.example.com
|
||||||
|
```
|
||||||
|
|
||||||
|
Updates are the primary lifecycle event on Cloudron and are designed to preserve `/app/data`. After the app comes back, confirm health and that your admin npub is still recognised:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl https://routstr.example.com/health
|
||||||
|
routstrd npubs list
|
||||||
|
```
|
||||||
|
|
||||||
|
## Day-to-day operations
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cloudron logs -f --app routstr.example.com # follow both processes' output
|
||||||
|
cloudron exec --app routstr.example.com # shell into the container
|
||||||
|
cloudron stop --app routstr.example.com
|
||||||
|
cloudron start --app routstr.example.com
|
||||||
|
cloudron debug --app routstr.example.com # read-write filesystem, app paused
|
||||||
|
cloudron debug --disable --app routstr.example.com
|
||||||
|
```
|
||||||
|
|
||||||
|
Inside the container you are on the machine that holds the wallet and the shared database, so `routstrd` commands there operate on the node itself rather than as a remote member:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd npubs list # everyone with access, and their roles
|
||||||
|
routstrd clients list # every client on the node, not just your own
|
||||||
|
routstrd top # interactive usage TUI across all members
|
||||||
|
```
|
||||||
|
|
||||||
|
Both processes log to stdout/stderr and are collected by Cloudron; there are no log files to rotate inside `/app/data`.
|
||||||
|
|
||||||
|
### Failure handling
|
||||||
|
|
||||||
|
`supervisord` runs both programs with `autorestart=true` and a start priority that brings the **daemon up first** (priority 10) and the **proxy second** (priority 20). The proxy's launcher additionally waits — up to 120 seconds — for the database file to exist *and* for the daemon's `/health` to answer before it starts serving. A crash-looping proxy therefore usually means the daemon never became healthy.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Next steps
|
||||||
|
|
||||||
|
- [Team Members](team-members.md) — invite the rest of your team.
|
||||||
|
- [Security Model](security.md) — what is exposed and what is not.
|
||||||
|
- [Troubleshooting](troubleshooting.md) — when the proxy will not start.
|
||||||
@@ -0,0 +1,170 @@
|
|||||||
|
# Deploy with Docker
|
||||||
|
|
||||||
|
The team node is a single container running two supervised processes. You can run that same image on any Docker host and terminate TLS with your own reverse proxy.
|
||||||
|
|
||||||
|
!!! note "Read this first"
|
||||||
|
The image is built `FROM cloudron/base:5.0.0` and its entrypoint is the Cloudron startup script. It works outside Cloudron — it only requires a writable `/app/data` — but its filesystem conventions are Cloudron's, and the repository's `docker-compose.yml` predates this image. The commands below are the ones that match the current image.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Build
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git clone https://github.com/routstr/routstrd-remote
|
||||||
|
cd routstrd-remote
|
||||||
|
docker build -t routstr-remote:0.1.26 .
|
||||||
|
```
|
||||||
|
|
||||||
|
The Dockerfile installs Bun (the x64 **baseline** build, chosen because some hosts do not expose AVX/AVX2 and the default binary crashes with `SIGILL`), installs the `routstrd` daemon globally, installs the proxy's dependencies, and copies in the supervisor configuration for both processes.
|
||||||
|
|
||||||
|
## Run
|
||||||
|
|
||||||
|
```bash
|
||||||
|
mkdir -p "$HOME/routstr-remote-data"
|
||||||
|
|
||||||
|
docker run -d \
|
||||||
|
--name routstr-remote \
|
||||||
|
--restart unless-stopped \
|
||||||
|
-p 127.0.0.1:8008:8008 \
|
||||||
|
-v "$HOME/routstr-remote-data:/app/data" \
|
||||||
|
--memory 1g \
|
||||||
|
routstr-remote:0.1.26
|
||||||
|
```
|
||||||
|
|
||||||
|
**Why each flag matters:**
|
||||||
|
|
||||||
|
| Flag | Reason |
|
||||||
|
|---|---|
|
||||||
|
| `-v ...:/app/data` | The only persistent path. It holds `routstrd/config.json` (including the container's `nsec`), `routstrd/routstr.db`, and `logs/`. **Without it the node loses its identity and every npub on restart.** |
|
||||||
|
| `-p 127.0.0.1:8008:8008` | Publish on loopback only and let a TLS reverse proxy in front of it expose the service. Binding `0.0.0.0:8008` on a public host sends API keys over plaintext HTTP. |
|
||||||
|
| `--memory 1g` | Two Bun processes plus the daemon's model and usage state. The Cloudron manifest requests 512 MB; give a bare Docker host at least that, and prefer more. |
|
||||||
|
|
||||||
|
The container listens on exactly two ports: `8008` (public auth proxy) and `8009` (daemon, bound to loopback **inside** the container and deliberately not published).
|
||||||
|
|
||||||
|
Check it came up:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://127.0.0.1:8008/health
|
||||||
|
docker logs -f routstr-remote
|
||||||
|
```
|
||||||
|
|
||||||
|
### Startup ordering
|
||||||
|
|
||||||
|
The entrypoint starts `supervisord`, which brings up the daemon first and the proxy second. The proxy's launcher then waits for the database file to appear and for `http://localhost:8009/health` to answer, retrying up to 120 times at one-second intervals. If that window expires, the proxy exits with `Timed out waiting for routstrd to become ready.` and restarts.
|
||||||
|
|
||||||
|
## Put TLS in front
|
||||||
|
|
||||||
|
Anything that terminates TLS and forwards to `127.0.0.1:8008` works. Two properties are worth configuring explicitly:
|
||||||
|
|
||||||
|
- **Disable response buffering.** The proxy already sends `X-Accel-Buffering: no` upstream, but your own proxy should also be configured not to buffer, otherwise streamed LLM responses appear to truncate.
|
||||||
|
- **Raise idle timeouts.** Model responses can be silent for a long time while reasoning or waiting on tools. The proxy disables Bun's per-request idle timeout for exactly this reason, but nginx's default 60-second `proxy_read_timeout` will still cut streams. Set it to something generous.
|
||||||
|
|
||||||
|
```nginx
|
||||||
|
location / {
|
||||||
|
proxy_pass http://127.0.0.1:8008;
|
||||||
|
proxy_http_version 1.1;
|
||||||
|
proxy_set_header Host $host;
|
||||||
|
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||||
|
proxy_set_header X-Forwarded-Proto $scheme;
|
||||||
|
|
||||||
|
proxy_buffering off;
|
||||||
|
proxy_read_timeout 3600s;
|
||||||
|
proxy_send_timeout 3600s;
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Bootstrap the first admin
|
||||||
|
|
||||||
|
Identical to Cloudron — run this from the admin's own machine:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bun i -g routstrd
|
||||||
|
routstrd remote https://routstr.example.com
|
||||||
|
routstrd npubs register --name "Alice"
|
||||||
|
```
|
||||||
|
|
||||||
|
While the npub table is empty, `POST /npubs` needs no authentication, so the first registration wins. See [Team Members](team-members.md).
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
Override any of the variables from [Deploy on Cloudron](deploy-cloudron.md#configuration) with `-e`, for example:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker run -d \
|
||||||
|
--name routstr-remote \
|
||||||
|
--restart unless-stopped \
|
||||||
|
-p 127.0.0.1:8008:8008 \
|
||||||
|
-v "$HOME/routstr-remote-data:/app/data" \
|
||||||
|
-e ROUTSTRD_AUTH_MODEL_ALLOWLIST=true \
|
||||||
|
routstr-remote:0.1.26
|
||||||
|
```
|
||||||
|
|
||||||
|
Inside the container the defaults are `ROUTSTRD_DIR=/app/data/routstrd`, `ROUTSTRD_DB_PATH=/app/data/routstrd/routstr.db`, `ROUTSTRD_UPSTREAM=http://localhost:8009`, `ROUTSTRD_AUTH_HOST=0.0.0.0`, `ROUTSTRD_AUTH_PORT=8008`.
|
||||||
|
|
||||||
|
!!! warning "Use a bind mount, not an anonymous volume"
|
||||||
|
If you recreate the container (`docker rm` then `docker run`), an unnamed volume is orphaned and the node comes back with a **new** Nostr identity and an **empty** npub table — meaning the next person to hit `POST /npubs` becomes admin. Always mount a known host directory.
|
||||||
|
|
||||||
|
## Backups
|
||||||
|
|
||||||
|
The database is SQLite in WAL mode, so stop the container before copying the data directory, or take a proper snapshot from inside:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker stop routstr-remote
|
||||||
|
tar czf routstr-remote-$(date +%F).tar.gz -C "$HOME" routstr-remote-data
|
||||||
|
docker start routstr-remote
|
||||||
|
```
|
||||||
|
|
||||||
|
Or, without downtime:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker exec routstr-remote sqlite3 /app/data/routstrd/routstr.db ".backup /app/data/routstrd/backup.db"
|
||||||
|
docker cp routstr-remote:/app/data/routstrd/backup.db .
|
||||||
|
```
|
||||||
|
|
||||||
|
## Updates
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker stop routstr-remote && docker rm routstr-remote
|
||||||
|
git pull
|
||||||
|
docker build -t routstr-remote:new .
|
||||||
|
docker run -d ... routstr-remote:new # same -v and -p flags as before
|
||||||
|
```
|
||||||
|
|
||||||
|
Because `/app/data` is a bind mount, the node's identity, npub table, and client records survive the swap. That is the entire reason the mount is not optional.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Running from source (evaluation and development)
|
||||||
|
|
||||||
|
You do not need Docker to try the proxy. Point it at any routstrd daemon's database:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bun install
|
||||||
|
bun run src/index.ts validate # checks config and opens the DB
|
||||||
|
bun run src/index.ts start # binds 0.0.0.0:8008
|
||||||
|
```
|
||||||
|
|
||||||
|
Useful flags: `--port`, `--host`, `--upstream`, `--db-path`. The `validate` subcommand prints the effective configuration and reports how many npubs are registered, split by role — it is the fastest way to confirm the proxy can see the right database:
|
||||||
|
|
||||||
|
```text
|
||||||
|
Configuration:
|
||||||
|
Port: 8008
|
||||||
|
Host: 0.0.0.0
|
||||||
|
Upstream: http://localhost:8009
|
||||||
|
DB path: /app/data/routstrd/routstr.db
|
||||||
|
Bootstrap admin npubs/pubkeys from env: 0
|
||||||
|
Model allowlist: disabled
|
||||||
|
|
||||||
|
✅ DB accessible. 3 npub(s) registered (1 admin, 2 user).
|
||||||
|
```
|
||||||
|
|
||||||
|
!!! warning "`validate` fails before the daemon has run once"
|
||||||
|
If the database does not exist, validation stops with `Database not found at ... Make sure routstrd has been initialized`. The proxy shares the daemon's database; it never creates the schema itself.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Next steps
|
||||||
|
|
||||||
|
- [Team Members](team-members.md) — invite the rest of your team.
|
||||||
|
- [Security Model](security.md) — which endpoints require what.
|
||||||
|
- [Troubleshooting](troubleshooting.md) — startup and streaming failures.
|
||||||
@@ -0,0 +1,112 @@
|
|||||||
|
# Teams and Remote Nodes
|
||||||
|
|
||||||
|
A normal `routstrd` install is a **single-user daemon** running on your own machine. A **remote node** turns that same daemon into a **shared instance for a team**: one server runs the daemon behind an authentication proxy, and every team member gets their own Nostr identity, their own API keys, and their own usage accounting.
|
||||||
|
|
||||||
|
This is the product sometimes called **Routstrd Remote** — the repository is [`Routstr/routstrd-remote`](https://github.com/routstr/routstrd-remote) and the running service is the package `routstrd-auth`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Who this is for
|
||||||
|
|
||||||
|
- **Small teams and orgs** that want one funded Routstr endpoint instead of one daemon per laptop.
|
||||||
|
- **Anyone who wants per-person spend visibility** without building a billing system.
|
||||||
|
- **Self-hosters** who already run [Cloudron](https://www.cloudron.io/) and want a one-click install.
|
||||||
|
|
||||||
|
## What it gives you
|
||||||
|
|
||||||
|
| Capability | Detail |
|
||||||
|
|---|---|
|
||||||
|
| **One endpoint, many people** | Members point their coding agents at a single HTTPS URL. No per-machine setup beyond one CLI command. |
|
||||||
|
| **Per-member identity** | Each person has their own Nostr keypair (npub). Access is granted and revoked by adding or deleting that npub. |
|
||||||
|
| **Per-member attribution** | Every client registration gets a unique ID, and usage is reported per client, so you can see who is spending what. |
|
||||||
|
| **Two levels of privilege** | `admin` (manage people, move funds, control the node) and `user` (run inference, manage only their own clients). |
|
||||||
|
| **Scoped API keys** | Agent API keys can buy inference but cannot touch the wallet or other members' clients. |
|
||||||
|
| **Model policy** | Optional allowlist restricts the team to approved models. |
|
||||||
|
| **One wallet to fund** | The team tops up a single node wallet rather than N personal wallets. |
|
||||||
|
|
||||||
|
## What it is not
|
||||||
|
|
||||||
|
- It is **not** multi-tenant SaaS. Everyone shares the node's wallet and upstream provider set.
|
||||||
|
- It is **not** a billing or chargeback system. Usage is *attributed* per client; invoicing your teammates is up to you.
|
||||||
|
- It does **not** give each member a separate balance. See [Usage and Model Policy](usage-and-policy.md) for exactly what is tracked.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Architecture
|
||||||
|
|
||||||
|
Two processes run together inside one container. Only one of them is reachable from outside.
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
flowchart TD
|
||||||
|
CLI["routstrd CLI<br/>on a member laptop"]
|
||||||
|
Agent["Coding agents<br/>Claude Code, Pi, OpenCode"]
|
||||||
|
App["App holding an sk- API key"]
|
||||||
|
|
||||||
|
TLS["Reverse proxy<br/>TLS termination on 443"]
|
||||||
|
Proxy["routstrd-auth<br/>0.0.0.0:8008 public"]
|
||||||
|
Daemon["routstrd daemon<br/>localhost:8009 no auth"]
|
||||||
|
DB[("routstr.db<br/>shared SQLite")]
|
||||||
|
Providers["Upstream model providers"]
|
||||||
|
|
||||||
|
CLI -->|https| TLS
|
||||||
|
Agent -->|https| TLS
|
||||||
|
App -->|https| TLS
|
||||||
|
TLS -->|http| Proxy
|
||||||
|
Proxy -->|forward| Daemon
|
||||||
|
Proxy -->|npubs and clients| DB
|
||||||
|
Daemon -->|usage and models| DB
|
||||||
|
Daemon -->|inference| Providers
|
||||||
|
```
|
||||||
|
|
||||||
|
**The security property that matters:** the daemon runs with **no authentication at all** because it is bound to `localhost` and never published. The auth proxy is the only public surface. If you expose port `8009`, you have removed the entire security model.
|
||||||
|
|
||||||
|
### Components
|
||||||
|
|
||||||
|
| Component | Role | Bind |
|
||||||
|
|---|---|---|
|
||||||
|
| `routstrd-auth` | Public auth proxy. Validates credentials, enforces roles and model policy, forwards to the daemon. | `0.0.0.0:8008` |
|
||||||
|
| `routstrd` | The inference daemon. Owns the wallet, providers, clients, and usage records. | `localhost:8009` |
|
||||||
|
| `routstr.db` | Shared SQLite database. Holds `routstr_auth_npubs`, `clients`, usage rows, and `sdk_storage` (including the Routstr 21 model list). | on disk |
|
||||||
|
| Reverse proxy | TLS termination and the public hostname. On Cloudron this is managed by the platform. | `443` |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Roles
|
||||||
|
|
||||||
|
Registration lives in a single table, `routstr_auth_npubs`, where each row has a `role` of `admin` or `user`.
|
||||||
|
|
||||||
|
| Capability | `admin` | `user` |
|
||||||
|
|---|---|---|
|
||||||
|
| Run inference with own API keys | yes | yes |
|
||||||
|
| Create and delete **own** clients | yes | yes |
|
||||||
|
| Read **own** usage | yes | yes |
|
||||||
|
| List all registered npubs | yes | yes |
|
||||||
|
| Add / update / delete npubs | yes | no |
|
||||||
|
| Send funds from the node wallet | yes | no |
|
||||||
|
| Node control (providers, refunds, stop) | yes | yes |
|
||||||
|
| Read wallet balance / status | yes | yes |
|
||||||
|
|
||||||
|
`user` is the default role when an admin adds someone. Promote with `routstrd npubs update <npub> --role admin`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Bootstrap order
|
||||||
|
|
||||||
|
A fresh node has an empty npub table. That produces exactly one unaudited window, and only one:
|
||||||
|
|
||||||
|
1. **The first person** runs `routstrd npubs register` against the new node. Because the table is empty, `POST /npubs` is accepted **without authentication**, and the caller becomes `admin`.
|
||||||
|
2. From that moment on, **every** npub operation requires NIP-98 auth from an existing admin. A second unauthenticated registration is refused with `409` / "already configured".
|
||||||
|
|
||||||
|
Nobody else can self-register. Team members must be added by an admin. See [Team Members](team-members.md).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Next steps
|
||||||
|
|
||||||
|
- **[Deploy on Cloudron](deploy-cloudron.md)** — the supported, packaged deployment with TLS and backups handled for you.
|
||||||
|
- **[Deploy with Docker](deploy-docker.md)** — run the same image anywhere, behind your own reverse proxy.
|
||||||
|
- **[Team Members](team-members.md)** — bootstrap the first admin and invite people.
|
||||||
|
- **[Connecting Clients](clients.md)** — wire up Claude Code, Pi, OpenCode, and raw API keys.
|
||||||
|
- **[Usage and Model Policy](usage-and-policy.md)** — per-member spend tracking and the model allowlist.
|
||||||
|
- **[Security Model](security.md)** — the exact auth rules, public paths, and restricted endpoints.
|
||||||
|
- **[Troubleshooting](troubleshooting.md)** — diagnosing the failures people actually hit.
|
||||||
@@ -0,0 +1,172 @@
|
|||||||
|
# Security Model
|
||||||
|
|
||||||
|
The whole design rests on one property: **the daemon has no authentication because it is never reachable.** `routstrd` binds loopback-only on port `8009`; the auth proxy on `8008` is the single public surface, and it is the only component that makes authorisation decisions.
|
||||||
|
|
||||||
|
!!! danger "Never publish port 8009"
|
||||||
|
The daemon is unauthenticated by design. Exposing it — or port-forwarding it for debugging, or forgetting to restrict a Docker port mapping to loopback — removes the entire security model at once. Anyone who reaches it can read the wallet, list every member's clients, and spend the team's funds.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Request decision flow
|
||||||
|
|
||||||
|
The proxy is **default-deny**. A request is only forwarded if some rule explicitly allows it.
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
flowchart TD
|
||||||
|
A["request arrives on 8008"] --> B{"management path?<br/>npubs, clients, usage"}
|
||||||
|
B -->|yes| C["own handler<br/>NIP-98 required"]
|
||||||
|
B -->|no| D{"GET or HEAD<br/>on a public path?"}
|
||||||
|
D -->|yes| E["forward, no auth"]
|
||||||
|
D -->|no| F{"Authorization header?"}
|
||||||
|
F -->|missing| G["401"]
|
||||||
|
F -->|"Bearer sk-..."| H{"key found in clients?"}
|
||||||
|
H -->|no| I["401"]
|
||||||
|
H -->|yes| J{"restricted path?<br/>wallet, node control"}
|
||||||
|
J -->|yes| K["403"]
|
||||||
|
J -->|no| L["forward with header intact"]
|
||||||
|
F -->|"Nostr event"| M{"valid NIP-98?<br/>url, method, body hash, sig"}
|
||||||
|
M -->|no| N["401"]
|
||||||
|
M -->|yes| O{"pubkey registered?<br/>and role sufficient?"}
|
||||||
|
O -->|no| P["403"]
|
||||||
|
O -->|yes| Q["forward, header stripped"]
|
||||||
|
```
|
||||||
|
|
||||||
|
Two details are easy to miss:
|
||||||
|
|
||||||
|
- **Public means `GET`/`HEAD` only.** The proxy applies the public-path rule only to read methods, because the daemon routes a `POST` to those same paths as a *paid* request. `GET /v1/models` needs no credential; `POST /v1/models` does.
|
||||||
|
- **Management paths are matched before anything else.** `/npubs`, `/clients`, `/clients/add`, `/clients/delete`, `/usage`, and `/usage/summary` are handled by the proxy's own handlers and never forwarded to the daemon wholesale.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Public paths
|
||||||
|
|
||||||
|
Reachable with no credential at all (`GET`/`HEAD` only):
|
||||||
|
|
||||||
|
| Path | Purpose |
|
||||||
|
|---|---|
|
||||||
|
| `/health` | Liveness. Used by Cloudron's health check and by the proxy's own upstream probe. |
|
||||||
|
| `/ping` | Lightweight reachability. |
|
||||||
|
| `/models` | Model directory. |
|
||||||
|
| `/v1/models` | OpenAI-compatible model list. |
|
||||||
|
| `/models/*`, `/v1/models/*` | Prefixes covering per-model detail paths. |
|
||||||
|
|
||||||
|
This is intentional — an agent needs to discover models before it has a key, and provider discovery is public information in Routstr.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Credential one: API keys (`Bearer sk-...`)
|
||||||
|
|
||||||
|
An API key is looked up in the client records. If no client carries it, the request is rejected with `401 Invalid API key.`
|
||||||
|
|
||||||
|
Keys are **deliberately narrow**. A valid key is refused with `403` on every restricted path:
|
||||||
|
|
||||||
|
| Restricted endpoint | Why |
|
||||||
|
|---|---|
|
||||||
|
| `/wallet/status`, `/wallet/unlock`, `/wallet/balance` | An inference key must not read wallet state. |
|
||||||
|
| `/wallet/receive/cashu`, `/wallet/receive/bolt11` | No minting funds with a key. |
|
||||||
|
| `/wallet/send/cashu`, `/wallet/send/bolt11` | Admin-only in any case. |
|
||||||
|
| `/wallet/mints`, `/wallet/mints/info` | No mint inspection. |
|
||||||
|
| `/stop`, `/refund`, `/refund/xcashu` | No node control. |
|
||||||
|
| `/providers`, `/providers/enable`, `/providers/disable` | `?refresh=true` rewrites the stored provider list. |
|
||||||
|
| `/nwc/*` | No payment-channel access. |
|
||||||
|
| `/npubs`, `/clients/add`, `/clients/delete`, `/usage` | No management surface at all. |
|
||||||
|
|
||||||
|
The rule of thumb: **an API key buys inference, and nothing else.**
|
||||||
|
|
||||||
|
Key handling on the way through: the `Authorization` header is **preserved** so the daemon can validate the key itself, and its own accounting stays authoritative.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Credential two: NIP-98 (`Nostr <base64-event>`)
|
||||||
|
|
||||||
|
Management and wallet operations require a signed [NIP-98](https://github.com/nostr-protocol/nips/blob/master/98.md) event. This is not a bearer token — it is a signature over the specific request, so it cannot be replayed against a different endpoint.
|
||||||
|
|
||||||
|
The proxy enforces, in order:
|
||||||
|
|
||||||
|
| Check | Rule |
|
||||||
|
|---|---|
|
||||||
|
| Event kind | must be `27235` |
|
||||||
|
| Timestamp | within **±60 seconds** of now |
|
||||||
|
| `u` tag | must equal the **absolute request URL**, including scheme and host |
|
||||||
|
| `method` tag | must match the HTTP method (case-insensitive) |
|
||||||
|
| `payload` tag | **required when the body is non-empty**; must equal the SHA-256 hex digest of the raw body, compared in constant time |
|
||||||
|
| Signature | verified with `verifyEvent` |
|
||||||
|
|
||||||
|
The proxy then looks the pubkey up in `routstr_auth_npubs`:
|
||||||
|
|
||||||
|
- **Not registered** → `403`. The error message is context-aware: on a node with no npubs at all it tells you to run `routstrd npubs register`; otherwise it says registered auth is required.
|
||||||
|
- **Registered but role insufficient** → `403 Admin access required.`
|
||||||
|
- **Registered and sufficient** → forwarded, with the `Authorization` header **stripped** so it does not reach the daemon or the upstream provider.
|
||||||
|
|
||||||
|
!!! warning "Behind a reverse proxy, forwarded headers are not optional"
|
||||||
|
The `u` tag is checked against the **public** URL the client signed. The proxy reconstructs that URL from `X-Forwarded-Proto` and `X-Forwarded-Host` (falling back to `Host`). If your reverse proxy does not set them, the comparison fails and every NIP-98 request is rejected with `NIP-98 URL tag does not match this request.` See the nginx snippet in [Deploy with Docker](deploy-docker.md#put-tls-in-front).
|
||||||
|
|
||||||
|
!!! note "The ±60 second window means clock skew matters"
|
||||||
|
A client whose clock is more than a minute off will produce events that are rejected as `outside the allowed window`. If one machine alone fails to authenticate, check its clock before suspecting the node.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Role requirements by endpoint
|
||||||
|
|
||||||
|
| Endpoint group | Required |
|
||||||
|
|---|---|
|
||||||
|
| `/wallet/send/cashu`, `/wallet/send/bolt11` | `admin` |
|
||||||
|
| `/wallet/status`, `/wallet/unlock`, `/wallet/balance`, `/wallet/receive/*`, `/wallet/mints*`, `/stop`, `/refund*`, `/providers*`, `/nwc/*` | any registered npub (`admin` or `user`) |
|
||||||
|
| `/clients`, `/clients/add`, `/clients/delete` | any registered npub, **scoped to own clients** |
|
||||||
|
| `/usage`, `/usage/summary` | any registered npub, **scoped to own usage** |
|
||||||
|
| `/npubs` read | any registered npub |
|
||||||
|
| `/npubs` create / update / delete | `admin` |
|
||||||
|
| Everything else | valid API key **or** registered npub |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Bootstrap window
|
||||||
|
|
||||||
|
While `routstr_auth_npubs` is empty, `POST /npubs` is accepted **without authentication**. This exists solely so a fresh node can be claimed, and it closes permanently after the first registration.
|
||||||
|
|
||||||
|
The practical implication: a node that is deployed and healthy but has not had its first admin register is **unclaimed**. Treat deployment and bootstrap as one operation.
|
||||||
|
|
||||||
|
There is deliberately no hardcoded default admin in the image. The absence of one means an image cannot be shipped with a known admin key — but it also means a half-finished deployment is claimable by whoever finds it first.
|
||||||
|
|
||||||
|
As an alternative to the interactive step, admins can be seeded with `ROUTSTRD_AUTH_ADMIN_NPUBS`, `ROUTSTRD_AUTH_ADMIN_PUBKEYS`, or `ROUTSTRD_AUTH_BOOTSTRAP_NPUB`. Rows created this way are tagged `source = 'env'` and **reconciled at every startup** — remove the value from the environment and the row is deleted, which is a clean way to authorise a node declaratively.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## CORS
|
||||||
|
|
||||||
|
The proxy answers with:
|
||||||
|
|
||||||
|
```text
|
||||||
|
Access-Control-Allow-Origin: *
|
||||||
|
Access-Control-Allow-Methods: GET, POST, PATCH, DELETE, OPTIONS
|
||||||
|
Access-Control-Allow-Headers: Authorization, Content-Type, X-Cashu, X-Routstr-Model
|
||||||
|
Access-Control-Expose-Headers: X-Cashu, X-Routstr-Request-Id, X-Routstr-Cost-Msats,
|
||||||
|
X-Routstr-Cost-Usd, X-Routstr-Input-Cost-Msats,
|
||||||
|
X-Routstr-Output-Cost-Msats
|
||||||
|
```
|
||||||
|
|
||||||
|
A wildcard origin is safe **here specifically** because the app uses no cookies and no sessions. There is no ambient browser identity for a cross-origin page to borrow — a malicious page cannot make an authenticated request on a visitor's behalf, because every non-public request still needs its own API key or signature.
|
||||||
|
|
||||||
|
!!! warning "If you ever add cookie or session auth, revisit this"
|
||||||
|
The wildcard is only correct while authentication is entirely credential-based. Adding session cookies would turn this into a real vulnerability.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Hardening checklist
|
||||||
|
|
||||||
|
- [ ] **Daemon is loopback-only.** Verify `8009` is not published (`docker port routstr-remote`, or check the Cloudron app's port config).
|
||||||
|
- [ ] **TLS everywhere.** No member or agent should ever send an `sk-...` key over plaintext HTTP.
|
||||||
|
- [ ] **First admin registered** immediately after install.
|
||||||
|
- [ ] **Reverse proxy sets `X-Forwarded-Proto` / `X-Forwarded-Host`**, or NIP-98 fails.
|
||||||
|
- [ ] **Streaming hangs are fixed with timeouts, not by buffering.** Disable `proxy_buffering` and raise `proxy_read_timeout`; do not "fix" a truncated stream by publishing the daemon directly.
|
||||||
|
- [ ] **`ROUTSTRD_AUTH_ADMIN_NPUBS` reflects reality** if you use env bootstrapping — those rows are deleted on restart when the variable changes.
|
||||||
|
- [ ] **Departed members have their clients deleted**, not just their npub. Deleting an npub does **not** revoke existing API keys.
|
||||||
|
- [ ] **Backups cover `/app/data`** and are tested, including the `nsec` in `routstrd/config.json`.
|
||||||
|
- [ ] **Model allowlist verified** if you rely on it — it fails open when the model list has not been populated.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Next steps
|
||||||
|
|
||||||
|
- [Troubleshooting](troubleshooting.md) — diagnosing `401` and `403` responses.
|
||||||
|
- [Team Members](team-members.md#what-revocation-does-and-does-not-do) — the two-step offboarding that revocation alone does not cover.
|
||||||
@@ -0,0 +1,177 @@
|
|||||||
|
# Team Members
|
||||||
|
|
||||||
|
Access to a team node is an entry in one table: `routstr_auth_npubs`. Each row holds a Nostr pubkey, an optional display name, and a `role` of `admin` or `user`. Adding someone grants access; deleting their row revokes it.
|
||||||
|
|
||||||
|
There are no passwords, no invite links, and no email addresses. Identity is a Nostr keypair that each person generates on their own machine.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## The invite loop
|
||||||
|
|
||||||
|
A new member does the first two steps themselves; an existing admin does the third.
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
sequenceDiagram
|
||||||
|
participant M as New member
|
||||||
|
participant A as Existing admin
|
||||||
|
participant N as Team node
|
||||||
|
|
||||||
|
M->>M: install the routstrd CLI
|
||||||
|
M->>N: set the remote URL
|
||||||
|
N-->>M: generates keypair and prints npub
|
||||||
|
M->>A: send npub out of band
|
||||||
|
A->>N: add npub with role user
|
||||||
|
N-->>A: access confirmed
|
||||||
|
M->>N: add a client integration
|
||||||
|
N-->>M: API key issued
|
||||||
|
```
|
||||||
|
|
||||||
|
### 1. The member installs the CLI and connects
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bun i -g routstrd
|
||||||
|
routstrd remote https://team.example.com
|
||||||
|
```
|
||||||
|
|
||||||
|
`routstrd remote` writes the node URL into `~/.routstrd/config.json` and, **only if no Nostr identity exists yet**, generates one and prints the npub:
|
||||||
|
|
||||||
|
```text
|
||||||
|
Remote daemon URL set to: https://team.example.com
|
||||||
|
|
||||||
|
A new Nostr identity has been generated for remote authentication.
|
||||||
|
Your npub: npub1abc...xyz
|
||||||
|
You can view it in the config file at: /home/bob/.routstrd/config.json
|
||||||
|
```
|
||||||
|
|
||||||
|
If you already had an identity, it is reused and no npub is printed — run `routstrd remote` with no arguments to display the current node and identity.
|
||||||
|
|
||||||
|
!!! tip "Your npub is not a secret"
|
||||||
|
It is a public key and safe to paste into a team chat. The corresponding `nsec` lives in `~/.routstrd/config.json` and **is** a secret: it signs every management request. Never share it, and treat any host that has it as holding that member's credentials.
|
||||||
|
|
||||||
|
### 2. The member sends their npub to an admin
|
||||||
|
|
||||||
|
Out of band — chat, ticket, whatever. There is no self-service join.
|
||||||
|
|
||||||
|
### 3. An admin adds them
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd npubs add npub1abc...xyz --name "Bob"
|
||||||
|
```
|
||||||
|
|
||||||
|
The role defaults to `user`. The new member can now run inference and manage their own clients. To make them an admin, pass `--role admin` (or promote later).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Bootstrap: the first admin
|
||||||
|
|
||||||
|
A brand-new node has an empty table, which is a special case: `POST /npubs` is accepted **without authentication** so that someone can claim the node.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd remote https://team.example.com
|
||||||
|
routstrd npubs register --name "Alice"
|
||||||
|
```
|
||||||
|
|
||||||
|
`npubs register` is deliberately narrow — it refuses to do anything if any npub already exists:
|
||||||
|
|
||||||
|
```text
|
||||||
|
Admin npubs already configured (3). Ask your admin to add your npub.
|
||||||
|
Your npub: npub1...
|
||||||
|
```
|
||||||
|
|
||||||
|
So `register` only ever works once per node. After that, `npubs add` is the command, and it requires admin NIP-98 auth.
|
||||||
|
|
||||||
|
!!! warning "Claim the node during install"
|
||||||
|
Between the app becoming healthy and the first `npubs register`, the node is unclaimed — anybody who reaches the URL can become admin. See [Deploy on Cloudron](deploy-cloudron.md#bootstrap-the-first-admin).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Command reference
|
||||||
|
|
||||||
|
All of these talk to the auth proxy over NIP-98, so the caller must be a registered npub, and the mutating ones require the `admin` role.
|
||||||
|
|
||||||
|
| Command | Role needed | Notes |
|
||||||
|
|---|---|---|
|
||||||
|
| `routstrd npubs list` | any registered | Shows role and name for everyone, and marks your own row with `→ you`. |
|
||||||
|
| `routstrd npubs register` | none, **once** | Only succeeds while the table is empty. |
|
||||||
|
| `routstrd npubs add <npub>` | admin | Accepts `npub1...` or a 64-char hex pubkey. `--role admin\|user` (default `user`), `--name`. |
|
||||||
|
| `routstrd npubs update <npub>` | admin | `--role` and/or `--name`. Passing `--name ""` clears the name. |
|
||||||
|
| `routstrd npubs delete <npub>` | admin | Revokes access. |
|
||||||
|
|
||||||
|
Names are trimmed and capped at 64 characters.
|
||||||
|
|
||||||
|
`routstrd npubs list` output looks like this:
|
||||||
|
|
||||||
|
```text
|
||||||
|
Npubs (3):
|
||||||
|
- npub1qqq...4f2 [admin] "Alice" → you
|
||||||
|
- npub1xxx...9k7 [user] "Bob"
|
||||||
|
- npub1zzz...3md [user]
|
||||||
|
```
|
||||||
|
|
||||||
|
If your own npub is missing from the list, the CLI tells you whom to send it to:
|
||||||
|
|
||||||
|
```text
|
||||||
|
Your npub is not in the npub list. Ask an admin to add your npub:
|
||||||
|
npub1yyy...0pl
|
||||||
|
```
|
||||||
|
|
||||||
|
## Underlying HTTP API
|
||||||
|
|
||||||
|
The CLI is a thin wrapper over four endpoints on the auth proxy. Useful for scripting or a custom onboarding form.
|
||||||
|
|
||||||
|
| Method | Path | Auth | Body |
|
||||||
|
|---|---|---|---|
|
||||||
|
| `GET` | `/npubs` | any registered npub | — |
|
||||||
|
| `POST` | `/npubs` | none **if the table is empty**, otherwise admin | `{ "npub": "npub1..." }` or `{ "pubkey": "<hex>" }`, plus optional `role` and `name` |
|
||||||
|
| `PATCH` | `/npubs` | admin | `{ "npub": "npub1..." }` plus `role` and/or `name` (`name: null` clears it) |
|
||||||
|
| `DELETE` | `/npubs/<npub-or-pubkey>` | admin | — (also accepts `/npubs?npub=...`) |
|
||||||
|
|
||||||
|
Responses:
|
||||||
|
|
||||||
|
- `GET /npubs` returns `{ "npubs": [ { "npub": "...", "name": "...", "role": "admin" } ] }`.
|
||||||
|
- Adding a pubkey that is **already registered** returns `409` rather than silently succeeding; use `PATCH` to change an existing entry.
|
||||||
|
- Every mutation is performed by, and recorded against, the requesting admin.
|
||||||
|
|
||||||
|
!!! note "Revocation is immediate"
|
||||||
|
Roles and removals are read from the database on **every request** with no caching layer. Deleting an npub stops that member's management access on their next request. Their **API keys are a separate matter** — see below.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## What revocation does and does not do
|
||||||
|
|
||||||
|
Both halves of offboarding are separate, and only the first is available through the normal member-facing API.
|
||||||
|
|
||||||
|
**1. Management access — revoked by deleting the npub.**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd npubs delete npub1abc...xyz
|
||||||
|
```
|
||||||
|
|
||||||
|
Roles and rows are read from the database on **every request** with no caching layer, so the next NIP-98 request from that key is rejected immediately.
|
||||||
|
|
||||||
|
**2. Inference keys — a separate, manually managed thing.**
|
||||||
|
|
||||||
|
The Bearer path validates an `sk-...` key by looking it up in the client records and nothing else. Deleting an npub does **not** touch those rows, so an offboarded member's existing API keys keep working for inference until the client itself is deleted.
|
||||||
|
|
||||||
|
Here the proxy's scoping matters: `/clients`, `/clients/add`, and `/clients/delete` are **strictly owner-scoped** — the proxy filters and authorises by the calling npub. Being an `admin` does not grant access to a *colleague's* client records through those endpoints. So the realistic offboarding paths are:
|
||||||
|
|
||||||
|
- **Have the member delete their own clients** before you remove their npub: `routstrd clients list`, then `routstrd clients delete <id>`.
|
||||||
|
- **Or do it on the node itself.** Inside the container the CLI talks to the daemon directly on loopback, where there is no auth layer and therefore no ownership filter:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cloudron exec --app routstr.example.com
|
||||||
|
routstrd clients list # every client on the node, all owners
|
||||||
|
routstrd clients delete <id>
|
||||||
|
```
|
||||||
|
|
||||||
|
On the node, client IDs appear **with** their owner suffix (`my-laptop-4f2x9k7`) — see [Connecting Clients](clients.md). That suffix is exactly what tells you which member a client belongs to.
|
||||||
|
|
||||||
|
!!! danger "Removing an npub is not the same as revoking access"
|
||||||
|
If you skip step 2, a departed member's agents continue to consume the team's wallet. Always pair `npubs delete` with deleting their client records.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Next steps
|
||||||
|
|
||||||
|
- [Connecting Clients](clients.md) — get each member's agents talking to the node.
|
||||||
|
- [Usage and Model Policy](usage-and-policy.md) — see what each person is spending.
|
||||||
@@ -0,0 +1,161 @@
|
|||||||
|
# Troubleshooting
|
||||||
|
|
||||||
|
Most team-node problems are one of five things: the daemon never came up, the wrong database, a credential problem, a clock or proxy-header problem, or a networking/timeout issue around streaming.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## First triage
|
||||||
|
|
||||||
|
Run these in order. Together they answer "is the node up, does it have people, and is the proxy seeing the right database".
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 1. Is the public surface alive?
|
||||||
|
curl -sS https://team.example.com/health
|
||||||
|
|
||||||
|
# 2. Do both processes run, and what did they log on boot?
|
||||||
|
cloudron logs --app routstr.example.com | tail -50 # or: docker logs routstr-remote
|
||||||
|
|
||||||
|
# 3. Can the proxy see the database, and does it know your people?
|
||||||
|
cloudron exec --app routstr.example.com
|
||||||
|
routstrd-auth validate
|
||||||
|
```
|
||||||
|
|
||||||
|
Step 3 is the most informative. Its output ends with a line like `✅ DB accessible. 3 npub(s) registered (1 admin, 2 user).` — if that count is wrong, or the path is wrong, you have found your problem.
|
||||||
|
|
||||||
|
On startup the proxy also logs a summary, and warns loudly if the node is unclaimed:
|
||||||
|
|
||||||
|
```text
|
||||||
|
routstrd-auth proxy listening on http://0.0.0.0:8008
|
||||||
|
Upstream: http://localhost:8009
|
||||||
|
DB path: /app/data/routstrd/routstr.db
|
||||||
|
Registered npubs: 0
|
||||||
|
Model allowlist: disabled
|
||||||
|
Warning: no registered npub/pubkey. The first admin can be registered without auth using POST /npubs.
|
||||||
|
```
|
||||||
|
|
||||||
|
!!! warning "`Registered npubs: 0` on an established node is an emergency"
|
||||||
|
Either the database was lost (usually a missing volume mount) or the proxy is pointed at the wrong file. While the table is empty the node is claimable by anyone who reaches it. Fix it before anything else — see [Lost identity or empty npub list](#lost-identity-or-empty-npub-list).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Error reference
|
||||||
|
|
||||||
|
### Startup
|
||||||
|
|
||||||
|
| Symptom | Cause | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| `Timed out waiting for routstrd to become ready.` | The proxy's launcher waited 120 seconds for the database file to exist **and** for `http://localhost:8009/health` to answer. The daemon is unhealthy or crashing. | Read the daemon's log lines above this one. It is a daemon problem, not a proxy problem. |
|
||||||
|
| `Database not found at /app/data/routstrd/routstr.db. Make sure routstrd has been initialized (routstrd onboard).` | The proxy shares the daemon's database and never creates the schema. It ran before the daemon had ever initialised. | Start the daemon first (`routstrd start`, or let `supervisord` do it — priority 10 before 20). |
|
||||||
|
| `Invalid admin Nostr pubkey(s): <value>. Use npub or 64-char hex pubkeys.` | A bootstrap admin variable contains something unparseable. | Fix `ROUTSTRD_AUTH_ADMIN_NPUBS` / `_PUBKEYS` / `_BOOTSTRAP_NPUB`. Use `npub1...` or 64-char hex. |
|
||||||
|
| Container exits with `Illegal instruction` / `SIGILL` | Bun's default x64 build needs AVX/AVX2, which some hosts and VMs do not expose. | The shipped image already uses the **baseline** build for this reason. If you build your own, keep `BUN_TARGET=bun-linux-x64-baseline`. |
|
||||||
|
| Proxy crash-loops immediately after a config change | Validation fails, so `start` exits non-zero and `supervisord` restarts it. | Run `routstrd-auth validate` to see the message. |
|
||||||
|
|
||||||
|
### Authentication
|
||||||
|
|
||||||
|
| Response | Meaning | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| `401 Missing Authorization header. Use 'Authorization: Bearer sk-...' or 'Authorization: Nostr <base64-event>'.` | No credential sent, and the path is not public. Remember public paths are `GET`/`HEAD` only. | Send a credential, or use a read method on a public path. |
|
||||||
|
| `401 Invalid API key.` | The key is not in the client records. | Recover it with `routstrd clients add --name "<existing name>"`, which prints the existing key. |
|
||||||
|
| `403 API keys cannot access this endpoint. Use NIP-98 auth from a registered npub/pubkey.` | You used an `sk-...` key on a wallet, node-control, or management endpoint. | Use the CLI, which signs with NIP-98 automatically. |
|
||||||
|
| `403 Admin access required. Only admin npubs can perform this action.` | The npub is registered but holds the `user` role. | An admin promotes them: `routstrd npubs update <npub> --role admin`. |
|
||||||
|
| `403 This endpoint requires a registered npub/pubkey, but none is configured. Register the first admin with 'routstrd npubs register'.` | The npub table is empty. | Bootstrap the first admin. |
|
||||||
|
| `403 This endpoint requires NIP-98 auth from a registered npub/pubkey.` | The signature was valid but the pubkey is not in the table. | An admin adds it: `routstrd npubs add <npub>`. |
|
||||||
|
| `401 Invalid Authorization format. Expected 'Bearer sk-...' or 'Nostr <base64-event>'.` | The header used a different scheme or a typo'd prefix. | Fix the prefix. |
|
||||||
|
|
||||||
|
A useful diagnostic: the token **type** determines the error. A `403` naming a specific capability means you authenticated successfully and were then refused by policy. A `401` means you did not authenticate at all.
|
||||||
|
|
||||||
|
### NIP-98 signature rejections
|
||||||
|
|
||||||
|
These all return `401` with a precise reason.
|
||||||
|
|
||||||
|
| Message | Cause | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| `NIP-98 URL tag does not match this request.` | The most common one. The `u` tag holds the URL the client signed, and the proxy reconstructs the public URL from `X-Forwarded-Proto` / `X-Forwarded-Host` (falling back to `Host`). Behind TLS termination, a missing forwarded header makes the reconstructed URL `http://...` while the client signed `https://...`. | Configure the reverse proxy to set both headers. See [Deploy with Docker](deploy-docker.md#put-tls-in-front). |
|
||||||
|
| `NIP-98 event timestamp is outside the allowed window.` | Events must be within **±60 seconds**. Clock skew. | Sync the client's clock (`timedatectl`, NTP). If only one machine fails, it is that machine. |
|
||||||
|
| `NIP-98 payload tag is required for requests with a body.` \| `NIP-98 payload tag does not match the request body hash.` | The signed SHA-256 does not match the body that arrived — something rewrote the body in transit. | Check for a middleware, WAF, or forward proxy that re-encodes request bodies. |
|
||||||
|
| `NIP-98 method tag does not match this request.` | The event was signed for a different method (often a `GET` signature reused on a `POST`). | Sign per request; do not reuse events. |
|
||||||
|
| `Invalid NIP-98 event signature.` \| `Invalid NIP-98 event kind.` | Corrupted token, or not a kind `27235` event. | Regenerate the request with the CLI. |
|
||||||
|
| `Invalid NIP-98 token encoding.` \| `Invalid NIP-98 event JSON.` | The base64 payload is truncated — common when a long `Authorization` header is split or truncated by a client. | Check for a header-size limit in the client or proxy. |
|
||||||
|
|
||||||
|
### Client-side CLI messages
|
||||||
|
|
||||||
|
| Message | Meaning | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| `The daemon at <url> rejected this account.` then `Register/authorize this npub on the remote daemon first: <npub>` | The node does not recognise your npub. | Send the printed npub to an admin. |
|
||||||
|
| `No remote node is set up.` | No `daemonUrl` in `~/.routstrd/config.json`. | `routstrd remote https://team.example.com`. |
|
||||||
|
| `Your npub is not in the npub list. Ask an admin to add your npub:` | `npubs list` worked (so you *are* registered) but shows you as absent — typically a stale local identity after a config reset. | Re-run `routstrd remote` to display your current npub, and confirm with an admin which one is registered. |
|
||||||
|
| `Daemon is not running` | The local daemon is unreachable. | Only relevant when running from source; `routstrd start`. |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Common scenarios
|
||||||
|
|
||||||
|
### Streaming responses get cut off mid-answer
|
||||||
|
|
||||||
|
The symptom is a reply that stops abruptly — often after exactly the same number of seconds — with no error from the model.
|
||||||
|
|
||||||
|
This is **not** a node problem. It is response buffering or an idle timeout in a proxy in front of the node. Model turns can be silent for a long time while reasoning or waiting on tools, and intermediaries treat that silence as a dead connection.
|
||||||
|
|
||||||
|
Fix, in order of what actually bites:
|
||||||
|
|
||||||
|
1. **Disable response buffering** where the proxy talks to `8008`. The node already sends `X-Accel-Buffering: no`, but the fronting proxy must also be told (`proxy_buffering off` in nginx).
|
||||||
|
2. **Raise read timeouts** well above the default (`proxy_read_timeout 3600s`). nginx's 60-second default is the usual culprit.
|
||||||
|
3. **Do not work around it by exposing the daemon.** Publishing `8009` removes authentication entirely.
|
||||||
|
|
||||||
|
The proxy itself disables Bun's per-request idle timeout (`server.timeout(req, 0)`) precisely because valid streams can be quiet for a long time; anything still timing out is outside the node.
|
||||||
|
|
||||||
|
### Lost identity or empty npub list
|
||||||
|
|
||||||
|
If `routstrd npubs list` is empty after a restart, or `routstrd-auth validate` reports `0 npub(s)`, the node lost its data directory. In Docker this is almost always an **unmounted or changed volume**: `/app/data` is the only persistent path.
|
||||||
|
|
||||||
|
The damage is two-fold:
|
||||||
|
|
||||||
|
- The `nsec` in `routstrd/config.json` is gone, so the node has a **new** identity.
|
||||||
|
- The npub table is empty, so `POST /npubs` is **unauthenticated again** — anyone who reaches the URL can claim the node.
|
||||||
|
|
||||||
|
Recover in this order:
|
||||||
|
|
||||||
|
1. **Stop the app** so nobody claims it.
|
||||||
|
2. Restore `/app/data` from a backup, or re-mount the correct volume and restart.
|
||||||
|
3. If data is unrecoverable, accept the new identity and re-register the first admin — then have everyone whose npub was lost send theirs again, and recreate their clients (keys do not survive either).
|
||||||
|
|
||||||
|
### Model requests return `403` unexpectedly
|
||||||
|
|
||||||
|
Either the sender is genuinely outside the allowlist, or the allowlist is enforcing a **stale** list.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cloudron exec --app routstr.example.com
|
||||||
|
sqlite3 /app/data/routstrd/routstr.db \
|
||||||
|
"SELECT substr(value,1,120) FROM sdk_storage WHERE key = 'routstr21Models';"
|
||||||
|
```
|
||||||
|
|
||||||
|
- **No output:** the daemon never bootstrapped the list, so the proxy is failing **open** and allowing everything. The `403` is coming from somewhere else.
|
||||||
|
- **Output present but missing the model you want:** the list is stale. Refresh it with `routstrd clients --manual-refresh`, and check the scheduled job has not been disabled with `--disable-automatic-refresh`.
|
||||||
|
|
||||||
|
Remember the check applies to `POST`/`PUT`/`PATCH` bodies containing a `model` field only, and comparisons are case-sensitive.
|
||||||
|
|
||||||
|
### A member cannot authenticate at all
|
||||||
|
|
||||||
|
Work down this list:
|
||||||
|
|
||||||
|
1. **Are they registered?** `routstrd npubs list` as an admin.
|
||||||
|
2. **Is their npub the one you registered?** They run `routstrd remote` with no arguments to print the identity actually in their config. A machine with an old or regenerated `nsec` presents a different npub.
|
||||||
|
3. **Is their clock correct?** Off by more than 60 seconds means every NIP-98 event is rejected.
|
||||||
|
4. **Are forwarded headers set?** If *everyone* is failing with a URL tag mismatch, it is the reverse proxy, not the people.
|
||||||
|
5. **Do they have admin-requiring needs?** A `403 Admin access required` is a role problem, not an authentication problem.
|
||||||
|
|
||||||
|
### The node works but nobody can see usage
|
||||||
|
|
||||||
|
Usage is attributed per client, and `/usage` is scoped to the caller's npub. A member with no clients has no usage to show. Aggregate figures require running the CLI **on the node**, where the daemon is unauthenticated and unfiltered:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cloudron exec --app routstr.example.com
|
||||||
|
routstrd top
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Next steps
|
||||||
|
|
||||||
|
- [Security Model](security.md) — the rules behind the `401`s and `403`s.
|
||||||
|
- [Usage and Model Policy](usage-and-policy.md) — allowlist behaviour and its fail-open case.
|
||||||
@@ -0,0 +1,144 @@
|
|||||||
|
# Usage and Model Policy
|
||||||
|
|
||||||
|
Two things a team admin cares about: **who is spending what**, and **which models the team may use**. The first is always tracked; the second is an opt-in policy.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Usage tracking
|
||||||
|
|
||||||
|
Every request is attributed to the **client** that made it, and every client is owned by exactly one member. That chain is what produces per-person reporting.
|
||||||
|
|
||||||
|
### As a member
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd usage # your own usage summary
|
||||||
|
routstrd top # interactive TUI, alias for 'monitor'
|
||||||
|
routstrd balance # wallet balance and status
|
||||||
|
```
|
||||||
|
|
||||||
|
### As the operator, on the node
|
||||||
|
|
||||||
|
The most complete view is the TUI, run inside the container:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cloudron exec --app routstr.example.com
|
||||||
|
routstrd top
|
||||||
|
```
|
||||||
|
|
||||||
|
The **clients** tab shows individual usage, and because each client ID carries the **last 7 characters of its owner's npub**, you can tell teammates apart at a glance even when several of them named their client `my-laptop`:
|
||||||
|
|
||||||
|
```text
|
||||||
|
my-laptop-4f2x9k7 1,204 req $3.18
|
||||||
|
my-laptop-9k7qp2 881 req $2.07
|
||||||
|
ci-runner-3md1zz 412 req $0.94
|
||||||
|
```
|
||||||
|
|
||||||
|
Run from the node, this view is **not** filtered to any one member — that is the difference between the CLI on the box and the same CLI on a laptop.
|
||||||
|
|
||||||
|
### Over HTTP
|
||||||
|
|
||||||
|
| Method | Path | Auth | Scope |
|
||||||
|
|---|---|---|---|
|
||||||
|
| `GET` | `/usage` | NIP-98 from a registered npub | The **caller's** usage only |
|
||||||
|
| `GET` | `/usage/summary` | NIP-98 from a registered npub | The **caller's** usage only |
|
||||||
|
|
||||||
|
The proxy forces the scope: it takes the authenticated npub and sets the `npub` query parameter before forwarding to the daemon, so the daemon returns that member's records. There is no query parameter you can pass to widen it — a member cannot read a colleague's usage through the API. Aggregate visibility requires access to the node itself.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## The wallet
|
||||||
|
|
||||||
|
The team shares **one node wallet**. Funding it is a single action rather than N personal top-ups, and that is one of the main reasons to run a team node.
|
||||||
|
|
||||||
|
| Endpoint | Required role |
|
||||||
|
|---|---|
|
||||||
|
| `/wallet/status`, `/wallet/balance`, `/wallet/mints`, `/wallet/mints/info` | any registered npub |
|
||||||
|
| `/wallet/receive/cashu`, `/wallet/receive/bolt11` | any registered npub |
|
||||||
|
| `/wallet/send/cashu`, `/wallet/send/bolt11` | **admin only** |
|
||||||
|
|
||||||
|
Reading the balance and receiving funds are open to every registered member; **moving funds out is admin-only.** Note the consequence: members can see the team's total balance but cannot withdraw from it. API keys can do none of this.
|
||||||
|
|
||||||
|
Funding, mint management, and payment semantics are the daemon's domain and are documented in the provider and client guides rather than repeated here.
|
||||||
|
|
||||||
|
!!! warning "Usage is attributed, not isolated"
|
||||||
|
There is no per-member balance and no spending cap on this layer. A member's clients spend from the same wallet everyone else does. If you need hard limits, they are not enforced here — track the usage view and manage access accordingly.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Model allowlist
|
||||||
|
|
||||||
|
The proxy can restrict the team to the **Routstr 21 model list**. This is **disabled by default**; enable it with:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ROUTSTRD_AUTH_MODEL_ALLOWLIST=true
|
||||||
|
```
|
||||||
|
|
||||||
|
When enabled, a request naming a model outside the list is rejected with `403` **before it reaches the daemon**, so no tokens are spent. When disabled, every model passes through untouched.
|
||||||
|
|
||||||
|
### How it works
|
||||||
|
|
||||||
|
```mermaid
|
||||||
|
flowchart TD
|
||||||
|
A["routstrd CLI or agent"] --> B["routstrd-auth<br/>checks auth, then model"]
|
||||||
|
B -->|"allowed"| C["routstrd daemon"]
|
||||||
|
B -->|"403 not allowed"| A
|
||||||
|
C --> D["upstream provider"]
|
||||||
|
E["Nostr kind 38423<br/>Routstr 21 list"] -->|fetched by daemon| F[("sdk_storage<br/>routstr21Models")]
|
||||||
|
F -->|read on every request| B
|
||||||
|
```
|
||||||
|
|
||||||
|
1. The **daemon's** SDK fetches the Routstr 21 list from Nostr (kind `38423` events) and stores it in the shared SQLite database under the `sdk_storage` table, key `routstr21Models`.
|
||||||
|
2. The **proxy** reads that key from the same database. It has **zero Nostr dependency** — no relay connections, no keys of its own for this purpose.
|
||||||
|
3. The value is read fresh on every request, with no caching, so a list update takes effect immediately.
|
||||||
|
|
||||||
|
Because the proxy shares the daemon's database, this costs a single indexed key lookup — typically under a millisecond.
|
||||||
|
|
||||||
|
### What is and is not checked
|
||||||
|
|
||||||
|
| Request | Checked |
|
||||||
|
|---|---|
|
||||||
|
| `POST` / `PUT` / `PATCH` with a JSON body containing `model` | yes |
|
||||||
|
| `GET` requests, including `/models` and `/v1/models` | no — public paths, forwarded immediately |
|
||||||
|
| Management endpoints (`/npubs`, `/clients`, `/usage`) | no — routed to their own handlers before this check |
|
||||||
|
| Non-JSON bodies | no — no `model` field can be extracted |
|
||||||
|
| Requests with no `model` field | no — forwarded, and the upstream produces the error |
|
||||||
|
|
||||||
|
Model IDs are compared **case-sensitively**, matching how they are stored in the list.
|
||||||
|
|
||||||
|
### Fail-open behaviour
|
||||||
|
|
||||||
|
If `routstr21Models` is absent — typically because the daemon has not bootstrapped it yet — the proxy **fails open and allows every model**. This is deliberate: the alternative is that a fresh node blocks all traffic until Nostr bootstrapping completes. The trade-off is that an allowlist can be silently ineffective early in a node's life, so verify it after enabling:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# confirm the daemon has populated the list
|
||||||
|
cloudron exec --app routstr.example.com
|
||||||
|
sqlite3 /app/data/routstrd/routstr.db \
|
||||||
|
"SELECT substr(value,1,120) FROM sdk_storage WHERE key = 'routstr21Models';"
|
||||||
|
```
|
||||||
|
|
||||||
|
If that returns nothing, the allowlist is not yet meaningful.
|
||||||
|
|
||||||
|
### Performance note
|
||||||
|
|
||||||
|
When enforcement is enabled, the proxy buffers `POST`/`PUT`/`PATCH` request bodies so it can inspect the `model` field. That is a small latency cost on request upload. **Response streaming is unaffected** — SSE and LLM token streams pass through unbuffered, and the proxy explicitly tells intermediaries not to buffer. When the allowlist is disabled, the body is not buffered at all.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Keeping the model list current
|
||||||
|
|
||||||
|
The list is updated by the daemon, not the proxy:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
routstrd clients --manual-refresh # refresh now
|
||||||
|
routstrd clients --disable-automatic-refresh # stop the scheduled job
|
||||||
|
routstrd clients --enable-automatic-refresh # resume it
|
||||||
|
```
|
||||||
|
|
||||||
|
If the scheduled refresh job is disabled and nobody runs a manual refresh, the allowlist enforces a **stale** list — and would eventually stop matching newly approved models.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Next steps
|
||||||
|
|
||||||
|
- [Security Model](security.md) — the full endpoint-by-endpoint auth matrix.
|
||||||
|
- [Troubleshooting](troubleshooting.md) — including "why am I getting a 403".
|
||||||
@@ -197,9 +197,9 @@ Properties:
|
|||||||
|
|
||||||
This is the only architecture that preserves end-to-end encryption from the user to the PPQ/Tinfoil enclave while still letting Routstr mediate payment. The key requirement is that usage/cost metadata must be returned outside the encrypted body, ideally as a response header available before body streaming begins.
|
This is the only architecture that preserves end-to-end encryption from the user to the PPQ/Tinfoil enclave while still letting Routstr mediate payment. The key requirement is that usage/cost metadata must be returned outside the encrypted body, ideally as a response header available before body streaming begins.
|
||||||
|
|
||||||
## Current Routstr problem
|
## Original Routstr problem
|
||||||
|
|
||||||
The current EHBP implementation charges successful EHBP requests at `max_cost_for_model` because Routstr cannot decrypt the response body:
|
The original EHBP implementation charged successful EHBP requests at `max_cost_for_model` because Routstr could not decrypt the response body:
|
||||||
|
|
||||||
```text
|
```text
|
||||||
successful EHBP request -> charge full reserved max cost
|
successful EHBP request -> charge full reserved max cost
|
||||||
@@ -260,7 +260,13 @@ During local PPQ testing, PPQ responses included this CORS exposure header:
|
|||||||
Access-Control-Expose-Headers: Ehbp-Response-Nonce, X-Private-Usage-Metrics, X-Encrypted-Usage-Metrics, X-Tinfoil-Usage-Metrics
|
Access-Control-Expose-Headers: Ehbp-Response-Nonce, X-Private-Usage-Metrics, X-Encrypted-Usage-Metrics, X-Tinfoil-Usage-Metrics
|
||||||
```
|
```
|
||||||
|
|
||||||
However, the actual tested non-streaming response did not include any of these usage headers, even when `X-Tinfoil-Request-Usage-Metrics: true` was sent.
|
However, when this was tested against PPQ's `/private/` endpoint
|
||||||
|
(`private/gpt-oss-120b`) on 2026-06-21, the non-streaming response did not
|
||||||
|
include any of these usage headers, even when
|
||||||
|
`X-Tinfoil-Request-Usage-Metrics: true` was sent. That observation does *not*
|
||||||
|
hold for the direct Tinfoil enclave upstream that Routstr ships: see
|
||||||
|
[Usage metrics header format](#usage-metrics-header-format) below, where the
|
||||||
|
response header and the streaming trailer are both verified present.
|
||||||
|
|
||||||
The decrypted body did include normal OpenAI usage, but only the decrypting Tinfoil client can see that body.
|
The decrypted body did include normal OpenAI usage, but only the decrypting Tinfoil client can see that body.
|
||||||
|
|
||||||
@@ -369,7 +375,7 @@ Possible approaches:
|
|||||||
|
|
||||||
- PPQ private models are billed per actual input/output tokens.
|
- PPQ private models are billed per actual input/output tokens.
|
||||||
- Private model rates are available from `GET /v1/models?type=all`.
|
- Private model rates are available from `GET /v1/models?type=all`.
|
||||||
- Current Routstr EHBP billing at max cost is wrong for PPQ private models.
|
- Max-cost EHBP fallback is wrong for PPQ private models; current code releases/refunds when trusted usage metadata is absent.
|
||||||
- Direct Tinfoil integration inside Routstr would enable exact usage billing but would make Routstr see plaintext.
|
- Direct Tinfoil integration inside Routstr would enable exact usage billing but would make Routstr see plaintext.
|
||||||
- A blind EHBP relay preserves privacy but requires PPQ/Tinfoil to expose usage/cost in plaintext headers/trailers.
|
- A blind EHBP relay preserves privacy but requires PPQ/Tinfoil to expose usage/cost in plaintext headers/trailers.
|
||||||
- The preferred solution is to keep Routstr blind and have PPQ return billing metadata outside the encrypted body.
|
- The preferred solution is to keep Routstr blind and have PPQ return billing metadata outside the encrypted body.
|
||||||
@@ -385,7 +391,9 @@ and `routstr/upstream/ehbp.py`.
|
|||||||
- Base URL: `https://inference.tinfoil.sh`
|
- Base URL: `https://inference.tinfoil.sh`
|
||||||
- Fetches models from the public `GET /v1/models` endpoint (no auth needed).
|
- Fetches models from the public `GET /v1/models` endpoint (no auth needed).
|
||||||
- Parses Tinfoil's pricing (`inputTokenPricePer1M`, `outputTokenPricePer1M`,
|
- Parses Tinfoil's pricing (`inputTokenPricePer1M`, `outputTokenPricePer1M`,
|
||||||
`requestPrice`) into the standard `Model`/`Pricing` schema.
|
`cachedInputTokenPricePer1M`, `requestPrice`) into the standard
|
||||||
|
`Model`/`Pricing` schema. Cached reads use the cached rate when present,
|
||||||
|
otherwise the full input rate; cache writes always use the full input rate.
|
||||||
- `supports_ehbp = True` — acts as a blind EHBP relay.
|
- `supports_ehbp = True` — acts as a blind EHBP relay.
|
||||||
- `get_ehbp_forwarding_target()` returns a target that includes
|
- `get_ehbp_forwarding_target()` returns a target that includes
|
||||||
`X-Tinfoil-Request-Usage-Metrics: true`.
|
`X-Tinfoil-Request-Usage-Metrics: true`.
|
||||||
@@ -396,9 +404,12 @@ and `routstr/upstream/ehbp.py`.
|
|||||||
|
|
||||||
- `routstr/upstream/ehbp.py`:
|
- `routstr/upstream/ehbp.py`:
|
||||||
- `parse_tinfoil_usage_metrics()` parses
|
- `parse_tinfoil_usage_metrics()` parses
|
||||||
`prompt=N,completion=N[,total=N][,model=<name>]` into an OpenAI-style
|
`prompt=N,completion=N[,total=N][,cached_prompt_tokens=N,
|
||||||
usage dict. The `model` field (added in tinfoilsh/confidential-model-router
|
uncached_prompt_tokens=N][,model=<name>][,cost_usd=<usd>]` into an
|
||||||
PR #385) is extracted as a string.
|
OpenAI-style usage dict. Cache reads map to ``cache_read_input_tokens``
|
||||||
|
so ``calculate_cost`` can bill them at the cached rate. The ``model``
|
||||||
|
field (added in tinfoilsh/confidential-model-router PR #385) is extracted
|
||||||
|
as a string; ``cost_usd`` is parsed for logging only.
|
||||||
- `_resolve_ehbp_target_url()` overrides the forwarding URL with
|
- `_resolve_ehbp_target_url()` overrides the forwarding URL with
|
||||||
`X-Tinfoil-Enclave-Url` when the SDK sends it.
|
`X-Tinfoil-Enclave-Url` when the SDK sends it.
|
||||||
- `_strip_proxy_headers()` removes `X-Routstr-Model`,
|
- `_strip_proxy_headers()` removes `X-Routstr-Model`,
|
||||||
@@ -410,8 +421,9 @@ and `routstr/upstream/ehbp.py`.
|
|||||||
actual served model's pricing is used for cost calculation.
|
actual served model's pricing is used for cost calculation.
|
||||||
- `forward_ehbp_request()` (bearer auth): if `X-Tinfoil-Usage-Metrics` is
|
- `forward_ehbp_request()` (bearer auth): if `X-Tinfoil-Usage-Metrics` is
|
||||||
present in the response header, finalizes with `adjust_payment_for_tokens()`
|
present in the response header, finalizes with `adjust_payment_for_tokens()`
|
||||||
for exact billing; otherwise falls back to max-cost. Billing uses the
|
for exact billing; otherwise releases the reservation. The encrypted body
|
||||||
actual served model when it differs from the requested one.
|
cannot be estimated locally, and the authorization ceiling is not billed.
|
||||||
|
Billing uses the actual served model when it differs from the requested one.
|
||||||
- `forward_ehbp_x_cashu_request()`: if usage is available, computes the
|
- `forward_ehbp_x_cashu_request()`: if usage is available, computes the
|
||||||
refund from actual cost instead of max cost, using the actual served
|
refund from actual cost instead of max cost, using the actual served
|
||||||
model's pricing when applicable.
|
model's pricing when applicable.
|
||||||
@@ -425,10 +437,10 @@ and `routstr/upstream/ehbp.py`.
|
|||||||
|---|---|---|
|
|---|---|---|
|
||||||
| Bearer, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Exact token cost via `adjust_payment_for_tokens` |
|
| Bearer, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Exact token cost via `adjust_payment_for_tokens` |
|
||||||
| Bearer, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Exact token cost (h11 captures trailers) |
|
| Bearer, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Exact token cost (h11 captures trailers) |
|
||||||
| Bearer, no usage header/trailer | N/A | Max-cost fallback |
|
| Bearer, no usage header/trailer | N/A | Release reservation; zero charge |
|
||||||
| X-Cashu, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Refund = `redeemed - actual_cost` |
|
| X-Cashu, non-streaming | `X-Tinfoil-Usage-Metrics` response header | Refund = `redeemed - actual_cost` |
|
||||||
| X-Cashu, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Refund = `redeemed - actual_cost` (h11 captures trailers) |
|
| X-Cashu, streaming | `X-Tinfoil-Usage-Metrics` HTTP trailer | Refund = `redeemed - actual_cost` (h11 captures trailers) |
|
||||||
| X-Cashu, no usage header/trailer | N/A | Refund = `redeemed - max_cost` |
|
| X-Cashu, no usage header/trailer | N/A | Full refund |
|
||||||
|
|
||||||
### Cost response headers
|
### Cost response headers
|
||||||
|
|
||||||
@@ -438,13 +450,20 @@ Routstr returns cost info as response headers:
|
|||||||
|
|
||||||
| Header | Auth | Description |
|
| Header | Auth | Description |
|
||||||
|---|---|---|
|
|---|---|---|
|
||||||
| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Total msats charged for this request |
|
| `X-Routstr-Cost-Msats` | Bearer, X-Cashu | Settled msats debited for this request |
|
||||||
| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the charge |
|
| `X-Routstr-Computed-Cost-Msats` | Bearer, X-Cashu | Computed usage cost, emitted when it differs from the settled debit |
|
||||||
|
| `X-Routstr-Cost-Usd` | Bearer | USD equivalent of the computed usage |
|
||||||
| `X-Routstr-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens |
|
| `X-Routstr-Input-Cost-Msats` | Bearer, X-Cashu | msats attributed to input tokens |
|
||||||
| `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens |
|
| `X-Routstr-Output-Cost-Msats` | Bearer, X-Cashu | msats attributed to output tokens |
|
||||||
|
| `X-Routstr-Cache-Read-Msats` | Bearer, X-Cashu | msats attributed to cached input (cache reads) |
|
||||||
|
| `X-Routstr-Cache-Creation-Msats` | Bearer, X-Cashu | msats attributed to cache creation (0 for Tinfoil today) |
|
||||||
|
|
||||||
The client/Tinfoil SDK can read these headers from the HTTP response without
|
The client/Tinfoil SDK can read these headers from the HTTP response without
|
||||||
needing to decrypt the body.
|
needing to decrypt the body. A duplicate or rejected finalization can therefore
|
||||||
|
report a zero settled debit while preserving the non-zero computed usage cost.
|
||||||
|
Normal JSON responses use the same contract: `total_msats` and `charged_msats`
|
||||||
|
are settled values, while `computed_msats` retains the usage calculation when
|
||||||
|
it differs.
|
||||||
|
|
||||||
### Setup
|
### Setup
|
||||||
|
|
||||||
@@ -461,9 +480,23 @@ Tinfoil returns usage metrics in the `X-Tinfoil-Usage-Metrics` response header
|
|||||||
true` is sent. As of tinfoilsh/confidential-model-router PR #385, the format is:
|
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>
|
prompt=<prompt_tokens>,completion=<completion_tokens>,total=<total_tokens>[,cached_prompt_tokens=<n>,uncached_prompt_tokens=<n>][,model=<served_model>][,cost_usd=<usd>]
|
||||||
```
|
```
|
||||||
|
|
||||||
|
`prompt` is the inclusive prompt total; `cached_prompt_tokens` is the portion
|
||||||
|
already in Tinfoil's prefix cache and is billed at the model's
|
||||||
|
`cachedInputTokenPricePer1M` rate (or the full input rate when the model has
|
||||||
|
no cached rate). `cost_usd` is Tinfoil's own computed request cost and is
|
||||||
|
currently parsed for observability only — Routstr bills from token counts.
|
||||||
|
|
||||||
|
Note that the header/trailer value is not always a single occurrence: for
|
||||||
|
streaming responses the trailer is emitted twice, so a client that reads the
|
||||||
|
trailer directly may see the same `prompt=...,completion=...,...` string twice
|
||||||
|
in one field, comma-joined. Parsers must be tolerant of the duplicate rather
|
||||||
|
than assuming exactly one occurrence. `parse_tinfoil_usage_metrics()` is
|
||||||
|
unaffected: it assigns each `key=value` part as it walks the comma-separated
|
||||||
|
value, and both occurrences carry identical numbers.
|
||||||
|
|
||||||
The `model` field carries the actual model name served by the enclave.
|
The `model` field carries the actual model name served by the enclave.
|
||||||
Routstr uses this to:
|
Routstr uses this to:
|
||||||
|
|
||||||
@@ -490,4 +523,6 @@ back to the requested model's pricing.
|
|||||||
finalizers for bearer and X-Cashu requests. This provides actual-cost billing
|
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.
|
today, at the cost of full time-to-last-byte latency for streaming responses.
|
||||||
- Whether Tinfoil's `/v1/responses` endpoint also returns usage metrics
|
- Whether Tinfoil's `/v1/responses` endpoint also returns usage metrics
|
||||||
headers or trailers.
|
headers or trailers. Verified: yes — `/v1/responses` returns
|
||||||
|
`X-Tinfoil-Usage-Metrics` as a plaintext response header, with the same field
|
||||||
|
set as `/v1/chat/completions` (including `cost_usd`).
|
||||||
|
|||||||
@@ -1,45 +0,0 @@
|
|||||||
import json
|
|
||||||
import sys
|
|
||||||
|
|
||||||
import httpx
|
|
||||||
|
|
||||||
|
|
||||||
def create_child_keys(base_url: str, api_key: str, count: int = 3) -> list[str]:
|
|
||||||
headers = {"Authorization": f"Bearer {api_key}"}
|
|
||||||
|
|
||||||
print(f"Requesting {count} child keys from {base_url}...")
|
|
||||||
|
|
||||||
child_keys = []
|
|
||||||
|
|
||||||
for i in range(count):
|
|
||||||
try:
|
|
||||||
response = httpx.post(f"{base_url}/v1/balance/child-key", headers=headers)
|
|
||||||
if response.status_code == 200:
|
|
||||||
data = response.json()
|
|
||||||
child_keys.append(data["api_key"])
|
|
||||||
print(
|
|
||||||
f" [{i + 1}] Created: {data['api_key']} (Cost: {data['cost_msats']} msats)"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
print(f" [{i + 1}] Failed: {response.status_code} - {response.text}")
|
|
||||||
except Exception as e:
|
|
||||||
print(f" [{i + 1}] Error: {str(e)}")
|
|
||||||
|
|
||||||
return child_keys
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
if len(sys.argv) < 2:
|
|
||||||
print("Usage: python create_child_keys.py <api_key_or_cashu_token> [base_url]")
|
|
||||||
sys.exit(1)
|
|
||||||
|
|
||||||
auth_key = sys.argv[1]
|
|
||||||
base_url = sys.argv[2] if len(sys.argv) > 2 else "http://localhost:8000"
|
|
||||||
|
|
||||||
keys = create_child_keys(base_url, auth_key)
|
|
||||||
|
|
||||||
if keys:
|
|
||||||
print("\nSuccessfully created child keys:")
|
|
||||||
print(json.dumps(keys, indent=2))
|
|
||||||
else:
|
|
||||||
print("\nNo child keys were created.")
|
|
||||||
+1
-1
@@ -7,7 +7,7 @@ from openai import OpenAI
|
|||||||
client = OpenAI(
|
client = OpenAI(
|
||||||
api_key=os.environ.get("TOKEN"),
|
api_key=os.environ.get("TOKEN"),
|
||||||
base_url=os.environ.get("ONION_URL", "http://roustrjfsdgfiueghsklchg.onion/v1"),
|
base_url=os.environ.get("ONION_URL", "http://roustrjfsdgfiueghsklchg.onion/v1"),
|
||||||
http_client=httpx.Client(proxies="socks5://localhost:9050"),
|
http_client=httpx.Client(proxy="socks5://localhost:9050"),
|
||||||
)
|
)
|
||||||
|
|
||||||
print(
|
print(
|
||||||
|
|||||||
@@ -0,0 +1,58 @@
|
|||||||
|
"""add refunds table
|
||||||
|
|
||||||
|
Revision ID: 3a0fbd387f10
|
||||||
|
Revises: a3f1b6c204de
|
||||||
|
Create Date: 2026-09-16
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
import sqlmodel
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision = "3a0fbd387f10"
|
||||||
|
down_revision = "a3f1b6c204de"
|
||||||
|
branch_labels = None
|
||||||
|
depends_on = None
|
||||||
|
|
||||||
|
OPEN_STATUSES = "status IN ('pending', 'ambiguous')"
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"refunds",
|
||||||
|
sa.Column("id", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
|
sa.Column(
|
||||||
|
"api_key_hashed_key", sqlmodel.sql.sqltypes.AutoString(), nullable=False
|
||||||
|
),
|
||||||
|
sa.Column("method", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
|
sa.Column("destination", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||||
|
sa.Column("amount_msats", sa.Integer(), nullable=False),
|
||||||
|
sa.Column("unit", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
|
sa.Column("mint_url", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
|
sa.Column("status", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
|
sa.Column("quote_id", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||||
|
sa.Column("token", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||||
|
sa.Column("claimed_at", sa.Integer(), nullable=True),
|
||||||
|
sa.Column("created_at", sa.Integer(), nullable=False),
|
||||||
|
sa.Column("updated_at", sa.Integer(), nullable=False),
|
||||||
|
sa.ForeignKeyConstraint(["api_key_hashed_key"], ["api_keys.hashed_key"]),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index("ix_refunds_api_key_hashed_key", "refunds", ["api_key_hashed_key"])
|
||||||
|
op.create_index("ix_refunds_status", "refunds", ["status"])
|
||||||
|
op.create_index(
|
||||||
|
"ux_refunds_open_per_key",
|
||||||
|
"refunds",
|
||||||
|
["api_key_hashed_key"],
|
||||||
|
unique=True,
|
||||||
|
sqlite_where=sa.text(OPEN_STATUSES),
|
||||||
|
postgresql_where=sa.text(OPEN_STATUSES),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_index("ux_refunds_open_per_key", table_name="refunds")
|
||||||
|
op.drop_index("ix_refunds_status", table_name="refunds")
|
||||||
|
op.drop_index("ix_refunds_api_key_hashed_key", table_name="refunds")
|
||||||
|
op.drop_table("refunds")
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
"""add model_metadata to model_paths
|
||||||
|
|
||||||
|
Revision ID: a3f1b6c204de
|
||||||
|
Revises: e5a6b7c8d9f0
|
||||||
|
Create Date: 2026-09-15 21:50:00.000000
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision = "a3f1b6c204de"
|
||||||
|
down_revision = "e5a6b7c8d9f0"
|
||||||
|
branch_labels = None
|
||||||
|
depends_on = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.add_column(
|
||||||
|
"model_paths",
|
||||||
|
sa.Column("model_metadata", sa.Text(), nullable=False, server_default="{}"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_column("model_paths", "model_metadata")
|
||||||
@@ -0,0 +1,76 @@
|
|||||||
|
"""Remove child keys and balance limits.
|
||||||
|
|
||||||
|
Removes the child-key feature (parent_key_hash) and the balance-limit
|
||||||
|
machinery (balance_limit, balance_limit_reset, balance_limit_reset_date)
|
||||||
|
from api_keys, plus the balance_limit/balance_limit_reset pass-through on
|
||||||
|
lightning_invoices.
|
||||||
|
|
||||||
|
Data preservation: before dropping the columns, every child key is
|
||||||
|
converted into a standalone key by clearing parent_key_hash. Child keys
|
||||||
|
never hold their own balance (they always spent from their parent), so no
|
||||||
|
funds are lost: the parent keeps its full balance, and the former child
|
||||||
|
rows are preserved with their total_spent/total_requests history intact.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision = "e5a6b7c8d9f0"
|
||||||
|
down_revision = "b4f7a1c9d2e3"
|
||||||
|
branch_labels = None
|
||||||
|
depends_on = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
# Convert child keys into standalone keys before dropping the link.
|
||||||
|
# Their balance is always 0 (they spent from the parent), so this
|
||||||
|
# cannot strand any funds.
|
||||||
|
op.execute("UPDATE api_keys SET parent_key_hash = NULL")
|
||||||
|
|
||||||
|
with op.batch_alter_table("api_keys") as batch_op:
|
||||||
|
batch_op.drop_index("ix_api_keys_parent_key_hash")
|
||||||
|
batch_op.drop_column("parent_key_hash")
|
||||||
|
batch_op.drop_column("balance_limit")
|
||||||
|
batch_op.drop_column("balance_limit_reset")
|
||||||
|
batch_op.drop_column("balance_limit_reset_date")
|
||||||
|
|
||||||
|
with op.batch_alter_table("lightning_invoices") as batch_op:
|
||||||
|
batch_op.drop_column("balance_limit")
|
||||||
|
batch_op.drop_column("balance_limit_reset")
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
with op.batch_alter_table("lightning_invoices") as batch_op:
|
||||||
|
batch_op.add_column(sa.Column("balance_limit", sa.Integer(), nullable=True))
|
||||||
|
batch_op.add_column(
|
||||||
|
sa.Column("balance_limit_reset", sa.String(), nullable=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
with op.batch_alter_table("api_keys") as batch_op:
|
||||||
|
batch_op.add_column(
|
||||||
|
sa.Column("balance_limit_reset_date", sa.Integer(), nullable=True)
|
||||||
|
)
|
||||||
|
batch_op.add_column(
|
||||||
|
sa.Column(
|
||||||
|
"balance_limit_reset",
|
||||||
|
sa.String(),
|
||||||
|
nullable=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
batch_op.add_column(sa.Column("balance_limit", sa.Integer(), nullable=True))
|
||||||
|
batch_op.add_column(
|
||||||
|
sa.Column(
|
||||||
|
"parent_key_hash",
|
||||||
|
sa.String(),
|
||||||
|
nullable=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
batch_op.create_foreign_key(
|
||||||
|
"fk_api_keys_parent_key_hash",
|
||||||
|
"api_keys",
|
||||||
|
["parent_key_hash"],
|
||||||
|
["hashed_key"],
|
||||||
|
)
|
||||||
|
batch_op.create_index(
|
||||||
|
"ix_api_keys_parent_key_hash", ["parent_key_hash"], unique=False
|
||||||
|
)
|
||||||
@@ -91,6 +91,15 @@ nav:
|
|||||||
- Advanced Pricing: provider/advanced-pricing.md
|
- Advanced Pricing: provider/advanced-pricing.md
|
||||||
- Discovery: provider/discovery.md
|
- Discovery: provider/discovery.md
|
||||||
- Tor Support: provider/tor.md
|
- Tor Support: provider/tor.md
|
||||||
|
- Teams (Remote):
|
||||||
|
- Overview: teams/index.md
|
||||||
|
- Deploy on Cloudron: teams/deploy-cloudron.md
|
||||||
|
- Deploy with Docker: teams/deploy-docker.md
|
||||||
|
- Team Members: teams/team-members.md
|
||||||
|
- Connecting Clients: teams/clients.md
|
||||||
|
- Usage and Model Policy: teams/usage-and-policy.md
|
||||||
|
- Security Model: teams/security.md
|
||||||
|
- Troubleshooting: teams/troubleshooting.md
|
||||||
- API Reference:
|
- API Reference:
|
||||||
- Overview: api/overview.md
|
- Overview: api/overview.md
|
||||||
- Authentication: api/authentication.md
|
- Authentication: api/authentication.md
|
||||||
|
|||||||
+36
-7
@@ -1,27 +1,27 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "routstr"
|
name = "routstr"
|
||||||
version = "0.4.5"
|
version = "0.4.7"
|
||||||
description = "Payment proxy for your LLM endpoint using cashu and nostr."
|
description = "Payment proxy for your LLM endpoint using cashu and nostr."
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.11"
|
requires-python = ">=3.11"
|
||||||
|
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"fastapi[standard]>=0.115",
|
"fastapi[standard-no-fastapi-cloud-cli]>=0.141",
|
||||||
"aiosqlite>=0.20",
|
"aiosqlite>=0.20",
|
||||||
"sqlmodel>=0.0.24",
|
"sqlmodel>=0.0.42", # Python 3.14 deferred-annotation support
|
||||||
"httpx[socks]>=0.25.2",
|
"httpx[socks]>=0.28.1",
|
||||||
"h11>=0.14",
|
"h11>=0.16",
|
||||||
"greenlet>=3.2.1",
|
"greenlet>=3.2.1",
|
||||||
"alembic>=1.13",
|
"alembic>=1.13",
|
||||||
"python-json-logger>=2.0.0",
|
"python-json-logger>=2.0.0",
|
||||||
"cashu>=0.20",
|
"cashu>=0.20",
|
||||||
"marshmallow>=3.13,<4.0",
|
"marshmallow>=3.13,<4.0",
|
||||||
"websockets>=12.0",
|
"websockets>=12.0",
|
||||||
"nostr>=0.0.2",
|
"nostr-sdk>=0.45.1,<0.46",
|
||||||
"mdurl==0.1.2",
|
"mdurl==0.1.2",
|
||||||
"pillow>=10",
|
"pillow>=10",
|
||||||
"openai>=1.98.0",
|
"openai>=1.98.0",
|
||||||
"litellm>=1.55.0",
|
"litellm>=1.93.0,<1.94", # 1.93 is the first line supporting Python 3.14
|
||||||
]
|
]
|
||||||
|
|
||||||
[dependency-groups]
|
[dependency-groups]
|
||||||
@@ -87,3 +87,32 @@ disallow_untyped_decorators = true
|
|||||||
|
|
||||||
[tool.uv.sources]
|
[tool.uv.sources]
|
||||||
routstr = { workspace = true }
|
routstr = { workspace = true }
|
||||||
|
|
||||||
|
# Security floors. cashu 0.20.x caps httpx, h11, fastapi, cryptography,
|
||||||
|
# setuptools and wheel below their patched versions. routstr uses only cashu's
|
||||||
|
# wallet modules, not its server paths, so the caps are safe to lift here.
|
||||||
|
[tool.uv]
|
||||||
|
# coincurve 20's build config uses cmake.verbose, removed in
|
||||||
|
# scikit-build-core 0.10. This keeps its source build working on Python 3.14.
|
||||||
|
build-constraint-dependencies = ["scikit-build-core<0.10"]
|
||||||
|
override-dependencies = [
|
||||||
|
# httpx 0.28 (required by litellm 1.93) removed the `proxies` kwarg cashu
|
||||||
|
# passes on every mint call; routstr/cashu_compat.py restores it for cashu.
|
||||||
|
"httpx[socks]>=0.28.1,<1.0",
|
||||||
|
"importlib-metadata>=8.0.0,<9.0",
|
||||||
|
"h11>=0.16.0",
|
||||||
|
"fastapi[standard-no-fastapi-cloud-cli]>=0.141",
|
||||||
|
"cryptography>=49.0.0",
|
||||||
|
"setuptools>=83.0.0",
|
||||||
|
"wheel>=0.46.2",
|
||||||
|
]
|
||||||
|
|
||||||
|
# Transitive deps whose dependents allow the patched version but don't require
|
||||||
|
# it. Constraints raise the floor without bypassing any upstream pin.
|
||||||
|
constraint-dependencies = [
|
||||||
|
"starlette>=1.3.1",
|
||||||
|
"httpcore>=1.0.9", # 1.0.8 caps h11<0.15
|
||||||
|
# 1.76 is the first grpcio-tools release with CPython 3.14 wheels.
|
||||||
|
"grpcio>=1.76.0,<2.0.0",
|
||||||
|
"grpcio-tools>=1.76.0,<2.0.0",
|
||||||
|
]
|
||||||
|
|||||||
+76
-16
@@ -59,6 +59,23 @@ def calculate_model_cost_score(model: "Model") -> float:
|
|||||||
return total_cost
|
return total_cost
|
||||||
|
|
||||||
|
|
||||||
|
def calculate_model_reservation_score(model: "Model") -> float:
|
||||||
|
"""Context-based reservation ceiling for ranking same-model candidates.
|
||||||
|
|
||||||
|
The balance gate reserves on ``sats_pricing.max_cost``, which scales with
|
||||||
|
``context_length``. Ranking on the same ceiling keeps the advertised,
|
||||||
|
routed, and reserved candidate consistent. Lower is better.
|
||||||
|
"""
|
||||||
|
from .payment.models import _calculate_usd_max_costs
|
||||||
|
|
||||||
|
try:
|
||||||
|
_, _, max_cost = _calculate_usd_max_costs(model)
|
||||||
|
return float(max_cost)
|
||||||
|
except Exception:
|
||||||
|
# Pricing shape missing/invalid; per-token score keeps order deterministic
|
||||||
|
return calculate_model_cost_score(model)
|
||||||
|
|
||||||
|
|
||||||
def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
|
def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
|
||||||
"""Calculate a penalty multiplier for certain providers.
|
"""Calculate a penalty multiplier for certain providers.
|
||||||
|
|
||||||
@@ -118,9 +135,21 @@ def create_model_mappings(
|
|||||||
Returns:
|
Returns:
|
||||||
Tuple of (model_instances, provider_map, unique_models)
|
Tuple of (model_instances, provider_map, unique_models)
|
||||||
"""
|
"""
|
||||||
from .payment.models import _row_to_model
|
from .payment.models import _row_to_model, has_usable_pricing
|
||||||
from .upstream.helpers import resolve_model_alias
|
from .upstream.helpers import resolve_model_alias
|
||||||
|
|
||||||
|
def _unusable_price(model: "Model") -> bool:
|
||||||
|
"""A candidate may only route on rates a request can be billed against.
|
||||||
|
|
||||||
|
Mirrors the served-catalog backstop in ``list_models``: a negative or
|
||||||
|
non-finite rate is not a price, and the cost calculation cannot bill on
|
||||||
|
one, so every request on the model would be charged the full maximum
|
||||||
|
reservation instead. Applies to provider-discovered models as well as
|
||||||
|
persisted overrides — no override row need exist for a malformed price
|
||||||
|
to be built into the candidate map.
|
||||||
|
"""
|
||||||
|
return not has_usable_pricing(model.pricing)
|
||||||
|
|
||||||
candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
|
candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
|
||||||
unique_models: dict[str, "Model"] = {}
|
unique_models: dict[str, "Model"] = {}
|
||||||
unique_model_keys: dict[str, str] = {}
|
unique_model_keys: dict[str, str] = {}
|
||||||
@@ -208,12 +237,33 @@ def create_model_mappings(
|
|||||||
# Apply overrides only for this provider's model row.
|
# Apply overrides only for this provider's model row.
|
||||||
if model_key is not None and model_key in overrides_by_key:
|
if model_key is not None and model_key in overrides_by_key:
|
||||||
override_row, provider_fee = overrides_by_key[model_key]
|
override_row, provider_fee = overrides_by_key[model_key]
|
||||||
|
try:
|
||||||
model_to_use = _row_to_model(
|
model_to_use = _row_to_model(
|
||||||
override_row, apply_provider_fee=True, provider_fee=provider_fee
|
override_row, apply_provider_fee=True, provider_fee=provider_fee
|
||||||
)
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
# Stored pricing is JSON from whatever wrote the row, so
|
||||||
|
# converting it can raise. Doing that inside this loop let
|
||||||
|
# one such row unwind the whole map build: at boot the node
|
||||||
|
# came up routing nothing, and on a later refresh the map it
|
||||||
|
# already had went permanently stale. The sibling loop over
|
||||||
|
# override-only rows already skips and logs such a row.
|
||||||
|
logger.warning(
|
||||||
|
"Skipping invalid model override while building model mappings",
|
||||||
|
extra={
|
||||||
|
"model_id": model.id,
|
||||||
|
"upstream_provider_id": upstream_db_id,
|
||||||
|
"error": str(exc),
|
||||||
|
"error_type": type(exc).__name__,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
continue
|
||||||
else:
|
else:
|
||||||
model_to_use = model
|
model_to_use = model
|
||||||
|
|
||||||
|
if _unusable_price(model_to_use):
|
||||||
|
continue
|
||||||
|
|
||||||
forwarded_model_id = get_effective_forwarded_model_id(model_to_use)
|
forwarded_model_id = get_effective_forwarded_model_id(model_to_use)
|
||||||
|
|
||||||
# Get all aliases for this model
|
# Get all aliases for this model
|
||||||
@@ -280,6 +330,8 @@ def create_model_mappings(
|
|||||||
continue
|
continue
|
||||||
if not model_to_use.enabled:
|
if not model_to_use.enabled:
|
||||||
continue
|
continue
|
||||||
|
if _unusable_price(model_to_use):
|
||||||
|
continue
|
||||||
|
|
||||||
forwarded_model_id = get_effective_forwarded_model_id(model_to_use)
|
forwarded_model_id = get_effective_forwarded_model_id(model_to_use)
|
||||||
|
|
||||||
@@ -323,20 +375,24 @@ def create_model_mappings(
|
|||||||
def alias_priority(model: "Model", alias: str) -> int:
|
def alias_priority(model: "Model", alias: str) -> int:
|
||||||
"""Rank how strong the mapping of alias->model is.
|
"""Rank how strong the mapping of alias->model is.
|
||||||
|
|
||||||
An exact model ID is authoritative and must be cost-ranked against the
|
A provider that serves the requested ID directly is authoritative and
|
||||||
other exact matches before considering forwarded aliases. This keeps a
|
must be cost-ranked before providers that only reach it through a
|
||||||
provider-specific forwarded ID from shadowing a directly available,
|
forwarded alias, so a forwarded ID cannot shadow a directly available,
|
||||||
cheaper model with the requested ID.
|
cheaper model with the requested ID.
|
||||||
|
|
||||||
|
"Directly served" covers both the exact model ID and the same ID behind
|
||||||
|
a provider prefix (e.g. ``gpt-oss-120b`` on Tinfoil vs
|
||||||
|
``openai/gpt-oss-120b`` on OpenRouter). Both name the same model, so
|
||||||
|
they share the top tier and cost decides between them; otherwise the
|
||||||
|
provider whose catalog omits the org prefix would always win on ID
|
||||||
|
spelling regardless of price.
|
||||||
"""
|
"""
|
||||||
if model.id and model.id.lower() == alias:
|
model_base = get_base_model_id(model.id)
|
||||||
return 5
|
if (model.id and model.id.lower() == alias) or model_base.lower() == alias:
|
||||||
|
return 4
|
||||||
|
|
||||||
forwarded_model_id = get_effective_forwarded_model_id(model)
|
forwarded_model_id = get_effective_forwarded_model_id(model)
|
||||||
if forwarded_model_id and forwarded_model_id.lower() == alias:
|
if forwarded_model_id and forwarded_model_id.lower() == alias:
|
||||||
return 4
|
|
||||||
|
|
||||||
model_base = get_base_model_id(model.id)
|
|
||||||
if model_base == alias:
|
|
||||||
return 3
|
return 3
|
||||||
if model.canonical_slug:
|
if model.canonical_slug:
|
||||||
canonical_base = get_base_model_id(model.canonical_slug)
|
canonical_base = get_base_model_id(model.canonical_slug)
|
||||||
@@ -345,15 +401,19 @@ def create_model_mappings(
|
|||||||
return 1
|
return 1
|
||||||
|
|
||||||
for alias, items in candidates.items():
|
for alias, items in candidates.items():
|
||||||
# Sort key: (priority DESC, cost ASC)
|
# Sort key: (priority DESC, reservation ASC, cost ASC)
|
||||||
# Using negative cost for DESC sort overall to keep high priority first
|
# Using negative costs for DESC sort overall to keep high priority first
|
||||||
def sort_key(item: tuple["Model", "BaseUpstreamProvider"]) -> tuple[int, float]:
|
def sort_key(
|
||||||
|
item: tuple["Model", "BaseUpstreamProvider"],
|
||||||
|
) -> tuple[int, float, float]:
|
||||||
model, provider = item
|
model, provider = item
|
||||||
priority = alias_priority(model, alias)
|
priority = alias_priority(model, alias)
|
||||||
cost = calculate_model_cost_score(model)
|
|
||||||
penalty = get_provider_penalty(provider)
|
penalty = get_provider_penalty(provider)
|
||||||
adjusted_cost = cost * penalty
|
# Rank on the reservation ceiling the balance gate enforces, with
|
||||||
return (priority, -adjusted_cost)
|
# per-token typical-usage cost as tiebreaker
|
||||||
|
adjusted_reservation = calculate_model_reservation_score(model) * penalty
|
||||||
|
adjusted_cost = calculate_model_cost_score(model) * penalty
|
||||||
|
return (priority, -adjusted_reservation, -adjusted_cost)
|
||||||
|
|
||||||
items.sort(key=sort_key, reverse=True)
|
items.sort(key=sort_key, reverse=True)
|
||||||
|
|
||||||
|
|||||||
+407
-495
File diff suppressed because it is too large
Load Diff
+46
-403
@@ -1,16 +1,13 @@
|
|||||||
import asyncio
|
|
||||||
import hashlib
|
import hashlib
|
||||||
import time
|
|
||||||
from time import monotonic
|
|
||||||
from typing import Annotated, NoReturn
|
from typing import Annotated, NoReturn
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from sqlmodel import col, select, update
|
from sqlmodel import col, select
|
||||||
|
|
||||||
|
from . import refund
|
||||||
from .auth import (
|
from .auth import (
|
||||||
get_billing_key,
|
|
||||||
redemption_error_to_http_exception,
|
redemption_error_to_http_exception,
|
||||||
validate_bearer_key,
|
validate_bearer_key,
|
||||||
)
|
)
|
||||||
@@ -21,20 +18,13 @@ from .core.db import (
|
|||||||
get_session,
|
get_session,
|
||||||
release_stale_reservations,
|
release_stale_reservations,
|
||||||
)
|
)
|
||||||
from .core.db import (
|
|
||||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
|
||||||
)
|
|
||||||
from .core.logging import get_logger
|
from .core.logging import get_logger
|
||||||
from .core.settings import settings
|
from .core.settings import settings
|
||||||
from .lightning import lightning_router
|
from .lightning import lightning_router
|
||||||
from .payment.lnurl import MeltOutcomeAmbiguousError
|
|
||||||
from .wallet import (
|
from .wallet import (
|
||||||
classify_redemption_error,
|
classify_redemption_error,
|
||||||
credit_balance,
|
credit_balance,
|
||||||
is_mint_connection_error,
|
|
||||||
recieve_token,
|
recieve_token,
|
||||||
send_to_lnurl,
|
|
||||||
send_token,
|
|
||||||
token_mint_url,
|
token_mint_url,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -58,39 +48,15 @@ async def get_key_from_header(
|
|||||||
|
|
||||||
|
|
||||||
async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict:
|
async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict:
|
||||||
billing_key = await get_billing_key(key, session)
|
|
||||||
info = {
|
info = {
|
||||||
"api_key": "sk-" + key.hashed_key,
|
"api_key": "sk-" + key.hashed_key,
|
||||||
"balance": billing_key.total_balance,
|
"balance": key.total_balance,
|
||||||
"reserved": billing_key.reserved_balance,
|
"reserved": key.reserved_balance,
|
||||||
"is_child": key.parent_key_hash is not None,
|
|
||||||
"total_requests": key.total_requests,
|
"total_requests": key.total_requests,
|
||||||
"total_spent": key.total_spent,
|
"total_spent": key.total_spent,
|
||||||
"balance_limit": key.balance_limit,
|
|
||||||
"balance_limit_reset": key.balance_limit_reset,
|
|
||||||
"validity_date": key.validity_date,
|
"validity_date": key.validity_date,
|
||||||
}
|
}
|
||||||
|
|
||||||
if key.parent_key_hash:
|
|
||||||
info["parent_key_preview"] = key.parent_key_hash[:8] + "..."
|
|
||||||
else:
|
|
||||||
# Fetch child keys if this is a parent key
|
|
||||||
statement = select(ApiKey).where(ApiKey.parent_key_hash == key.hashed_key)
|
|
||||||
results = await session.exec(statement)
|
|
||||||
child_keys = results.all()
|
|
||||||
if child_keys:
|
|
||||||
info["child_keys"] = [
|
|
||||||
{
|
|
||||||
"api_key": "sk-" + ck.hashed_key,
|
|
||||||
"total_requests": ck.total_requests,
|
|
||||||
"total_spent": ck.total_spent,
|
|
||||||
"balance_limit": ck.balance_limit,
|
|
||||||
"balance_limit_reset": ck.balance_limit_reset,
|
|
||||||
"validity_date": ck.validity_date,
|
|
||||||
}
|
|
||||||
for ck in child_keys
|
|
||||||
]
|
|
||||||
|
|
||||||
return info
|
return info
|
||||||
|
|
||||||
|
|
||||||
@@ -117,26 +83,18 @@ async def account_info(
|
|||||||
|
|
||||||
class BalanceCreateRequest(BaseModel):
|
class BalanceCreateRequest(BaseModel):
|
||||||
initial_balance_token: str
|
initial_balance_token: str
|
||||||
balance_limit: int | None = None
|
|
||||||
balance_limit_reset: str | None = None
|
|
||||||
validity_date: int | None = None
|
validity_date: int | None = None
|
||||||
|
|
||||||
|
|
||||||
async def _create_balance(
|
async def _create_balance(
|
||||||
initial_balance_token: str,
|
initial_balance_token: str,
|
||||||
balance_limit: int | None,
|
|
||||||
balance_limit_reset: str | None,
|
|
||||||
validity_date: int | None,
|
validity_date: int | None,
|
||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
key = await validate_bearer_key(initial_balance_token, session)
|
key = await validate_bearer_key(initial_balance_token, session)
|
||||||
|
|
||||||
if balance_limit is not None or balance_limit_reset or validity_date:
|
if validity_date is not None:
|
||||||
key.balance_limit = balance_limit
|
|
||||||
key.balance_limit_reset = balance_limit_reset
|
|
||||||
key.validity_date = validity_date
|
key.validity_date = validity_date
|
||||||
if balance_limit_reset:
|
|
||||||
key.balance_limit_reset_date = int(time.time())
|
|
||||||
session.add(key)
|
session.add(key)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
@@ -154,8 +112,6 @@ async def create_balance_from_body(
|
|||||||
) -> dict:
|
) -> dict:
|
||||||
return await _create_balance(
|
return await _create_balance(
|
||||||
payload.initial_balance_token,
|
payload.initial_balance_token,
|
||||||
payload.balance_limit,
|
|
||||||
payload.balance_limit_reset,
|
|
||||||
payload.validity_date,
|
payload.validity_date,
|
||||||
session,
|
session,
|
||||||
)
|
)
|
||||||
@@ -164,15 +120,11 @@ async def create_balance_from_body(
|
|||||||
@router.get("/create")
|
@router.get("/create")
|
||||||
async def create_balance(
|
async def create_balance(
|
||||||
initial_balance_token: str,
|
initial_balance_token: str,
|
||||||
balance_limit: int | None = None,
|
|
||||||
balance_limit_reset: str | None = None,
|
|
||||||
validity_date: int | None = None,
|
validity_date: int | None = None,
|
||||||
session: AsyncSession = Depends(get_session),
|
session: AsyncSession = Depends(get_session),
|
||||||
) -> dict:
|
) -> dict:
|
||||||
return await _create_balance(
|
return await _create_balance(
|
||||||
initial_balance_token,
|
initial_balance_token,
|
||||||
balance_limit,
|
|
||||||
balance_limit_reset,
|
|
||||||
validity_date,
|
validity_date,
|
||||||
session,
|
session,
|
||||||
)
|
)
|
||||||
@@ -208,7 +160,7 @@ async def topup_wallet_endpoint(
|
|||||||
key: ApiKey = Depends(get_key_from_header),
|
key: ApiKey = Depends(get_key_from_header),
|
||||||
session: AsyncSession = Depends(get_session),
|
session: AsyncSession = Depends(get_session),
|
||||||
) -> dict[str, int]:
|
) -> dict[str, int]:
|
||||||
billing_key = await get_billing_key(key, session)
|
billing_key = key
|
||||||
|
|
||||||
if topup_request is not None:
|
if topup_request is not None:
|
||||||
cashu_token = topup_request.cashu_token
|
cashu_token = topup_request.cashu_token
|
||||||
@@ -276,35 +228,6 @@ async def topup_wallet_endpoint(
|
|||||||
return {"msats": amount_msats}
|
return {"msats": amount_msats}
|
||||||
|
|
||||||
|
|
||||||
_REFUND_CACHE_TTL_SECONDS: int = settings.refund_cache_ttl_seconds
|
|
||||||
_refund_cache_lock: asyncio.Lock = asyncio.Lock()
|
|
||||||
_refund_cache: dict[str, tuple[float, dict[str, str]]] = {}
|
|
||||||
|
|
||||||
|
|
||||||
def _cache_key_for_authorization(authorization: str) -> str:
|
|
||||||
return hashlib.sha256(authorization.strip().encode()).hexdigest()
|
|
||||||
|
|
||||||
|
|
||||||
async def _refund_cache_get(authorization: str) -> dict[str, str] | None:
|
|
||||||
key = _cache_key_for_authorization(authorization)
|
|
||||||
async with _refund_cache_lock:
|
|
||||||
item = _refund_cache.get(key)
|
|
||||||
if item is None:
|
|
||||||
return None
|
|
||||||
expires_at, value = item
|
|
||||||
if expires_at <= monotonic():
|
|
||||||
del _refund_cache[key]
|
|
||||||
return None
|
|
||||||
return value
|
|
||||||
|
|
||||||
|
|
||||||
async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None:
|
|
||||||
key = _cache_key_for_authorization(authorization)
|
|
||||||
expiry = monotonic() + _REFUND_CACHE_TTL_SECONDS
|
|
||||||
async with _refund_cache_lock:
|
|
||||||
_refund_cache[key] = (expiry, value)
|
|
||||||
|
|
||||||
|
|
||||||
async def _lookup_key_no_create(
|
async def _lookup_key_no_create(
|
||||||
bearer_value: str, session: AsyncSession
|
bearer_value: str, session: AsyncSession
|
||||||
) -> ApiKey | None:
|
) -> ApiKey | None:
|
||||||
@@ -318,17 +241,16 @@ async def _lookup_key_no_create(
|
|||||||
|
|
||||||
|
|
||||||
async def _get_persisted_api_key_refund(
|
async def _get_persisted_api_key_refund(
|
||||||
key: ApiKey, session: AsyncSession
|
key: ApiKey, session: AsyncSession, token: str | None = None
|
||||||
) -> dict[str, str] | None:
|
) -> dict[str, str] | None:
|
||||||
result = await session.exec(
|
query = select(CashuTransaction).where(
|
||||||
select(CashuTransaction)
|
|
||||||
.where(
|
|
||||||
CashuTransaction.api_key_hashed_key == key.hashed_key,
|
CashuTransaction.api_key_hashed_key == key.hashed_key,
|
||||||
CashuTransaction.type == "out",
|
CashuTransaction.type == "out",
|
||||||
CashuTransaction.source == "apikey",
|
CashuTransaction.source == "apikey",
|
||||||
)
|
)
|
||||||
.order_by(col(CashuTransaction.created_at).desc())
|
if token is not None:
|
||||||
)
|
query = query.where(CashuTransaction.token == token)
|
||||||
|
result = await session.exec(query.order_by(col(CashuTransaction.created_at).desc()))
|
||||||
refund = result.first()
|
refund = result.first()
|
||||||
if refund is None:
|
if refund is None:
|
||||||
return None
|
return None
|
||||||
@@ -347,36 +269,13 @@ async def _get_persisted_api_key_refund(
|
|||||||
return persisted
|
return persisted
|
||||||
|
|
||||||
|
|
||||||
async def _restore_balance(
|
class RefundRequest(BaseModel):
|
||||||
session: AsyncSession,
|
lightning_address: str | None = None
|
||||||
hashed_key: str,
|
|
||||||
balance: int,
|
|
||||||
reserved_balance: int,
|
|
||||||
mint_url: str,
|
|
||||||
) -> None:
|
|
||||||
"""Restore balance after a failed refund mint attempt."""
|
|
||||||
restore_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == hashed_key)
|
|
||||||
.values(
|
|
||||||
balance=col(ApiKey.balance) + balance,
|
|
||||||
reserved_balance=col(ApiKey.reserved_balance) + reserved_balance,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
await session.exec(restore_stmt) # type: ignore[call-overload]
|
|
||||||
await session.commit()
|
|
||||||
logger.info(
|
|
||||||
"refund_wallet_endpoint: balance restored after mint failure",
|
|
||||||
extra={
|
|
||||||
"hashed_key": hashed_key,
|
|
||||||
"restored_balance": balance,
|
|
||||||
"mint_url": mint_url,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@router.post("/refund", response_model=None)
|
@router.post("/refund", response_model=None)
|
||||||
async def refund_wallet_endpoint(
|
async def refund_wallet_endpoint(
|
||||||
|
refund_request: RefundRequest | None = None,
|
||||||
authorization: Annotated[str | None, Header()] = None,
|
authorization: Annotated[str | None, Header()] = None,
|
||||||
x_cashu: Annotated[str | None, Header()] = None,
|
x_cashu: Annotated[str | None, Header()] = None,
|
||||||
session: AsyncSession = Depends(get_session),
|
session: AsyncSession = Depends(get_session),
|
||||||
@@ -452,17 +351,28 @@ async def refund_wallet_endpoint(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Check for an open claim before any replay or destination lookup.
|
||||||
|
if open_claim := await refund.latest_open(session, key):
|
||||||
|
raise refund.refund_in_progress_error(open_claim)
|
||||||
|
|
||||||
if key.total_balance <= 0:
|
if key.total_balance <= 0:
|
||||||
if cached := await _refund_cache_get(bearer_value):
|
paid = await refund.latest_terminal(session, key)
|
||||||
return cached
|
if paid and paid.method == "lightning":
|
||||||
|
return refund.describe(paid)
|
||||||
|
if paid and paid.token:
|
||||||
|
# Match the ledger row to this claim's token, not the latest one.
|
||||||
|
if persisted := await _get_persisted_api_key_refund(
|
||||||
|
key, session, paid.token
|
||||||
|
):
|
||||||
|
return persisted
|
||||||
|
return refund.describe(paid)
|
||||||
|
# Legacy payouts predate the claim row, so fall back to the ledger.
|
||||||
if persisted := await _get_persisted_api_key_refund(key, session):
|
if persisted := await _get_persisted_api_key_refund(key, session):
|
||||||
return persisted
|
return persisted
|
||||||
|
if paid:
|
||||||
if key.parent_key_hash:
|
return refund.describe(paid)
|
||||||
raise HTTPException(
|
if stuck := await refund.latest_stuck(session, key):
|
||||||
status_code=400,
|
raise refund.refund_in_progress_error(stuck)
|
||||||
detail="Cannot refund child key. Please refund the parent key instead.",
|
|
||||||
)
|
|
||||||
|
|
||||||
if key.reserved_balance > 0:
|
if key.reserved_balance > 0:
|
||||||
# Release only durable reservations old enough to be stale. A newer
|
# Release only durable reservations old enough to be stale. A newer
|
||||||
@@ -481,168 +391,33 @@ async def refund_wallet_endpoint(
|
|||||||
logger.warning(
|
logger.warning(
|
||||||
"refund_wallet_endpoint: released stale reservation before refund",
|
"refund_wallet_endpoint: released stale reservation before refund",
|
||||||
extra={
|
extra={
|
||||||
"hashed_key": key.hashed_key,
|
"key_hash": key.hashed_key[:8],
|
||||||
"stale_timeout_seconds": settings.stale_reservation_timeout_seconds,
|
"stale_timeout_seconds": settings.stale_reservation_timeout_seconds,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
remaining_balance_msats: int = key.total_balance
|
remaining_balance_msats: int = key.total_balance
|
||||||
|
unit = refund.refund_unit(key)
|
||||||
if key.refund_currency == "sat":
|
remaining_balance = refund.amount_in_unit(remaining_balance_msats, unit)
|
||||||
remaining_balance = remaining_balance_msats // 1000
|
|
||||||
else:
|
|
||||||
remaining_balance = remaining_balance_msats
|
|
||||||
|
|
||||||
if remaining_balance_msats > 0 and remaining_balance <= 0:
|
if remaining_balance_msats > 0 and remaining_balance <= 0:
|
||||||
raise HTTPException(status_code=400, detail="Balance too small to refund")
|
raise HTTPException(status_code=400, detail="Balance too small to refund")
|
||||||
elif remaining_balance <= 0:
|
elif remaining_balance <= 0:
|
||||||
raise HTTPException(status_code=400, detail="No balance to refund")
|
raise HTTPException(status_code=400, detail="No balance to refund")
|
||||||
|
|
||||||
# Capture values before debit — the session may refresh key after commit
|
requested = refund_request.lightning_address if refund_request else None
|
||||||
pre_debit_balance = key.balance
|
destination = requested or key.refund_address
|
||||||
pre_debit_reserved = key.reserved_balance
|
if destination:
|
||||||
|
# Stored addresses can rot too; reject before any balance is debited.
|
||||||
|
await refund.validate_lightning_destination(destination)
|
||||||
|
|
||||||
# --- DEBIT FIRST: atomically zero the balance before minting tokens ---
|
claim = await refund.open_claim(
|
||||||
# This prevents the race where a concurrent topup/spend happens between
|
|
||||||
# reading the balance and minting the refund token (double-spend).
|
|
||||||
debit_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.where(col(ApiKey.balance) == pre_debit_balance)
|
|
||||||
.where(col(ApiKey.reserved_balance) == pre_debit_reserved)
|
|
||||||
.values(balance=0, reserved_balance=0, reserved_at=None)
|
|
||||||
)
|
|
||||||
debit_result = await session.exec(debit_stmt) # type: ignore[call-overload]
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
if debit_result.rowcount == 0:
|
|
||||||
# Balance changed between read and debit — another request is active
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=409,
|
|
||||||
detail="Balance changed concurrently. Please retry the refund.",
|
|
||||||
)
|
|
||||||
|
|
||||||
# The balance is locked at zero, so it is safe to create the refund token.
|
|
||||||
effective_refund_mint = (
|
|
||||||
key.refund_mint_url
|
|
||||||
if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints
|
|
||||||
else settings.primary_mint
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
refund_currency = key.refund_currency or "sat"
|
|
||||||
if key.refund_address:
|
|
||||||
await send_to_lnurl(
|
|
||||||
remaining_balance,
|
|
||||||
key.refund_currency or "sat",
|
|
||||||
effective_refund_mint,
|
|
||||||
key.refund_address,
|
|
||||||
)
|
|
||||||
result = {"recipient": key.refund_address}
|
|
||||||
else:
|
|
||||||
token = await send_token(
|
|
||||||
remaining_balance, refund_currency, effective_refund_mint
|
|
||||||
)
|
|
||||||
effective_refund_mint = token_mint_url(token, effective_refund_mint)
|
|
||||||
result = {"token": token}
|
|
||||||
|
|
||||||
if key.refund_currency == "sat":
|
|
||||||
result["sats"] = str(remaining_balance_msats // 1000)
|
|
||||||
else:
|
|
||||||
result["msats"] = str(remaining_balance_msats)
|
|
||||||
|
|
||||||
if "token" in result:
|
|
||||||
logger.info(
|
|
||||||
"refund_wallet_endpoint: cashu token issued",
|
|
||||||
extra={
|
|
||||||
"path": "/v1/wallet/refund",
|
|
||||||
"token": result["token"],
|
|
||||||
"amount": remaining_balance,
|
|
||||||
"currency": key.refund_currency or "sat",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
except MeltOutcomeAmbiguousError as e:
|
|
||||||
# The melt was dispatched and may still settle. Restoring the balance
|
|
||||||
# here would let the same debit be paid out twice; keep the debit and
|
|
||||||
# leave the outcome to reconciliation.
|
|
||||||
logger.error(
|
|
||||||
"refund_wallet_endpoint: melt outcome ambiguous; balance withheld "
|
|
||||||
"pending reconciliation",
|
|
||||||
extra={
|
|
||||||
"error": str(e),
|
|
||||||
"hashed_key": key.hashed_key,
|
|
||||||
"remaining_balance": remaining_balance,
|
|
||||||
"refund_currency": key.refund_currency,
|
|
||||||
"refund_mint_url": key.refund_mint_url,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=502,
|
|
||||||
detail=(
|
|
||||||
"Refund was dispatched but its outcome is unconfirmed; the "
|
|
||||||
"balance is withheld until reconciliation completes"
|
|
||||||
),
|
|
||||||
)
|
|
||||||
except HTTPException:
|
|
||||||
# Minting failed — restore the debited balance
|
|
||||||
await _restore_balance(
|
|
||||||
session,
|
session,
|
||||||
key.hashed_key,
|
key,
|
||||||
pre_debit_balance,
|
method="lightning" if destination else "cashu",
|
||||||
pre_debit_reserved,
|
destination=destination,
|
||||||
key.refund_mint_url or "",
|
|
||||||
)
|
)
|
||||||
raise
|
return await refund.execute(session, claim)
|
||||||
except Exception as e:
|
|
||||||
# Minting failed — restore the debited balance
|
|
||||||
await _restore_balance(
|
|
||||||
session,
|
|
||||||
key.hashed_key,
|
|
||||||
pre_debit_balance,
|
|
||||||
pre_debit_reserved,
|
|
||||||
key.refund_mint_url or "",
|
|
||||||
)
|
|
||||||
error_msg = str(e)
|
|
||||||
logger.error(
|
|
||||||
"refund_wallet_endpoint: mint/send failed",
|
|
||||||
extra={
|
|
||||||
"error": error_msg,
|
|
||||||
"error_type": type(e).__name__,
|
|
||||||
"hashed_key": key.hashed_key,
|
|
||||||
"remaining_balance": remaining_balance,
|
|
||||||
"refund_currency": key.refund_currency,
|
|
||||||
"refund_mint_url": key.refund_mint_url,
|
|
||||||
"has_refund_address": bool(key.refund_address),
|
|
||||||
},
|
|
||||||
)
|
|
||||||
if is_mint_connection_error(e):
|
|
||||||
raise HTTPException(status_code=503, detail="Mint service unavailable")
|
|
||||||
else:
|
|
||||||
raise HTTPException(status_code=500, detail="Refund failed")
|
|
||||||
|
|
||||||
await _refund_cache_set(bearer_value, result)
|
|
||||||
|
|
||||||
if "token" in result:
|
|
||||||
await store_cashu_transaction(
|
|
||||||
token=result["token"],
|
|
||||||
amount=remaining_balance,
|
|
||||||
unit=key.refund_currency or "sat",
|
|
||||||
mint_url=effective_refund_mint,
|
|
||||||
typ="out",
|
|
||||||
collected=False,
|
|
||||||
source="apikey",
|
|
||||||
api_key_hashed_key=key.hashed_key,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"refund_wallet_endpoint: refund successful",
|
|
||||||
extra={
|
|
||||||
"refunded_msats": remaining_balance_msats,
|
|
||||||
"previous_reserved_balance": key.reserved_balance,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
@router.get("/history")
|
@router.get("/history")
|
||||||
@@ -650,12 +425,6 @@ async def wallet_history(
|
|||||||
key: ApiKey = Depends(get_key_from_header),
|
key: ApiKey = Depends(get_key_from_header),
|
||||||
session: AsyncSession = Depends(get_session),
|
session: AsyncSession = Depends(get_session),
|
||||||
) -> dict[str, list[dict[str, str | int | bool | None]]]:
|
) -> dict[str, list[dict[str, str | int | bool | None]]]:
|
||||||
if key.parent_key_hash:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=400,
|
|
||||||
detail="Cannot view child key history. Please use the parent key instead.",
|
|
||||||
)
|
|
||||||
|
|
||||||
result = await session.exec(
|
result = await session.exec(
|
||||||
select(CashuTransaction)
|
select(CashuTransaction)
|
||||||
.where(CashuTransaction.api_key_hashed_key == key.hashed_key)
|
.where(CashuTransaction.api_key_hashed_key == key.hashed_key)
|
||||||
@@ -693,132 +462,6 @@ async def donate(token: str, ref: str | None = None) -> str:
|
|||||||
return "Invalid token."
|
return "Invalid token."
|
||||||
|
|
||||||
|
|
||||||
class ChildKeyRequest(BaseModel):
|
|
||||||
count: int
|
|
||||||
balance_limit: int | None = None
|
|
||||||
balance_limit_reset: str | None = None
|
|
||||||
validity_date: int | None = None
|
|
||||||
|
|
||||||
|
|
||||||
@router.post("/child-key")
|
|
||||||
async def create_child_key(
|
|
||||||
payload: ChildKeyRequest,
|
|
||||||
key: ApiKey = Depends(get_key_from_header),
|
|
||||||
session: AsyncSession = Depends(get_session),
|
|
||||||
) -> dict:
|
|
||||||
"""Creates one or more child API keys that use the parent's balance."""
|
|
||||||
# Log incoming request for debugging
|
|
||||||
logger.debug(f"Child key creation request: count={payload.count}")
|
|
||||||
|
|
||||||
count = payload.count
|
|
||||||
if count < 1 or count > 50:
|
|
||||||
raise HTTPException(status_code=400, detail="Count must be between 1 and 50.")
|
|
||||||
|
|
||||||
# Check if this is already a child key
|
|
||||||
if key.parent_key_hash:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=400,
|
|
||||||
detail="Cannot create a child key for another child key.",
|
|
||||||
)
|
|
||||||
|
|
||||||
cost_per_key = settings.child_key_cost
|
|
||||||
total_cost = cost_per_key * count
|
|
||||||
|
|
||||||
if key.total_balance < total_cost:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=402,
|
|
||||||
detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Deduct cost from parent atomically — guards against concurrent requests
|
|
||||||
# that both pass the balance check above on stale in-memory state.
|
|
||||||
deduct_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.where(col(ApiKey.balance) - col(ApiKey.reserved_balance) >= total_cost)
|
|
||||||
.values(
|
|
||||||
balance=col(ApiKey.balance) - total_cost,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
result = await session.exec(deduct_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
if result.rowcount == 0:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=402,
|
|
||||||
detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Generate new keys
|
|
||||||
import secrets
|
|
||||||
|
|
||||||
new_keys = []
|
|
||||||
for _ in range(count):
|
|
||||||
new_key_raw = secrets.token_hex(32)
|
|
||||||
new_key_hash = new_key_raw # We use the raw key as the hash for sk- keys
|
|
||||||
|
|
||||||
child_key = ApiKey(
|
|
||||||
hashed_key=new_key_hash,
|
|
||||||
balance=0,
|
|
||||||
parent_key_hash=key.hashed_key,
|
|
||||||
balance_limit=payload.balance_limit,
|
|
||||||
balance_limit_reset=payload.balance_limit_reset,
|
|
||||||
balance_limit_reset_date=int(time.time())
|
|
||||||
if payload.balance_limit_reset
|
|
||||||
else None,
|
|
||||||
validity_date=payload.validity_date,
|
|
||||||
)
|
|
||||||
session.add(child_key)
|
|
||||||
new_keys.append("sk-" + new_key_hash)
|
|
||||||
|
|
||||||
await session.commit()
|
|
||||||
await session.refresh(key)
|
|
||||||
|
|
||||||
response_data = {
|
|
||||||
"api_keys": new_keys,
|
|
||||||
"count": count,
|
|
||||||
"cost_msats": total_cost,
|
|
||||||
"cost_sats": total_cost // 1000,
|
|
||||||
"parent_balance": key.balance,
|
|
||||||
"parent_balance_sats": key.balance // 1000,
|
|
||||||
}
|
|
||||||
logger.debug(f"Child key creation response: {response_data}")
|
|
||||||
return response_data
|
|
||||||
|
|
||||||
|
|
||||||
class ChildKeyResetRequest(BaseModel):
|
|
||||||
child_key: str
|
|
||||||
|
|
||||||
|
|
||||||
@router.post("/child-key/reset")
|
|
||||||
async def reset_child_key_spent(
|
|
||||||
payload: ChildKeyResetRequest,
|
|
||||||
key: ApiKey = Depends(get_key_from_header),
|
|
||||||
session: AsyncSession = Depends(get_session),
|
|
||||||
) -> dict:
|
|
||||||
"""Resets the total_spent of a child key. Must be called by the parent."""
|
|
||||||
child_key_raw = payload.child_key
|
|
||||||
if child_key_raw.startswith("sk-"):
|
|
||||||
child_key_raw = child_key_raw[3:]
|
|
||||||
|
|
||||||
child_key = await session.get(ApiKey, child_key_raw)
|
|
||||||
if not child_key:
|
|
||||||
raise HTTPException(status_code=404, detail="Child key not found.")
|
|
||||||
|
|
||||||
if child_key.parent_key_hash != key.hashed_key:
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=403, detail="Unauthorized. You are not the parent of this key."
|
|
||||||
)
|
|
||||||
|
|
||||||
child_key.total_spent = 0
|
|
||||||
if child_key.balance_limit_reset:
|
|
||||||
child_key.balance_limit_reset_date = int(time.time())
|
|
||||||
session.add(child_key)
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
return {"success": True, "message": "Child key balance reset successfully."}
|
|
||||||
|
|
||||||
|
|
||||||
@router.api_route(
|
@router.api_route(
|
||||||
"/{path:path}",
|
"/{path:path}",
|
||||||
methods=["GET", "POST", "PUT", "DELETE"],
|
methods=["GET", "POST", "PUT", "DELETE"],
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
"""Compatibility shim that keeps cashu 0.20.x working on httpx>=0.28.
|
||||||
|
|
||||||
|
cashu's ``async_set_httpx_client`` decorator builds the client for *every* mint
|
||||||
|
call as ``httpx.AsyncClient(proxies=proxies_dict, ...)``. httpx deprecated
|
||||||
|
``proxies`` in 0.26 and removed it in 0.28, so on httpx>=0.28 every wallet
|
||||||
|
operation routstr performs -- ``load_mint_keysets``, ``mint_quote``,
|
||||||
|
``melt_quote``, token redeem -- raises::
|
||||||
|
|
||||||
|
TypeError: AsyncClient.__init__() got an unexpected keyword argument 'proxies'
|
||||||
|
|
||||||
|
Upstream cashu (0.20.3, the latest release) still caps ``httpx<0.26`` and has no
|
||||||
|
release that fixes this, while litellm>=1.84 requires ``httpx>=0.28``. We bridge
|
||||||
|
the gap by translating the keyword *inside cashu's module namespace only* --
|
||||||
|
global httpx behaviour is untouched, and the shim is a no-op on httpx<0.28.
|
||||||
|
|
||||||
|
Delete this module once cashu ships a release that passes ``proxy=``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
__all__ = ["install_cashu_httpx_shim"]
|
||||||
|
|
||||||
|
_ALL_SCHEMES = "all://"
|
||||||
|
|
||||||
|
|
||||||
|
def _single_proxy(proxies: Any) -> str | None:
|
||||||
|
"""Collapse an httpx<0.28 ``proxies`` mapping into a single ``proxy`` URL.
|
||||||
|
|
||||||
|
cashu only ever builds ``{}`` or ``{"all://": url}``, so a mapping with one
|
||||||
|
distinct URL is all we need to support. Anything richer is unrepresentable
|
||||||
|
as httpx 0.28's scalar ``proxy=``; we fail closed and raise rather than
|
||||||
|
return ``None``, because dropping the entry would silently send mint
|
||||||
|
traffic direct instead of through the configured Tor/SOCKS proxy.
|
||||||
|
"""
|
||||||
|
if not proxies:
|
||||||
|
return None
|
||||||
|
if isinstance(proxies, str):
|
||||||
|
return proxies
|
||||||
|
if isinstance(proxies, dict):
|
||||||
|
if _ALL_SCHEMES in proxies:
|
||||||
|
value = proxies[_ALL_SCHEMES]
|
||||||
|
return str(value) if value is not None else None
|
||||||
|
distinct = {str(v) for v in proxies.values() if v is not None}
|
||||||
|
if not distinct:
|
||||||
|
return None
|
||||||
|
if len(distinct) == 1:
|
||||||
|
return distinct.pop()
|
||||||
|
raise ValueError(
|
||||||
|
f"cannot represent proxies={proxies!r} as httpx 0.28 proxy=; "
|
||||||
|
"refusing to send proxied traffic direct"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _ProxiesCompatAsyncClient(httpx.AsyncClient):
|
||||||
|
"""``httpx.AsyncClient`` that still accepts the removed ``proxies`` kwarg."""
|
||||||
|
|
||||||
|
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
||||||
|
if "proxies" in kwargs:
|
||||||
|
proxies = kwargs.pop("proxies")
|
||||||
|
proxy = _single_proxy(proxies)
|
||||||
|
if proxy is not None and kwargs.get("proxy") is None:
|
||||||
|
kwargs["proxy"] = proxy
|
||||||
|
elif isinstance(proxies, dict) and not proxies:
|
||||||
|
# In httpx<0.28, proxies={} disabled environment proxy discovery.
|
||||||
|
kwargs.setdefault("trust_env", False)
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
class _HttpxNamespace:
|
||||||
|
"""Stand-in for the ``httpx`` module inside cashu's ``v1_api``.
|
||||||
|
|
||||||
|
Every attribute resolves against the real module except ``AsyncClient``,
|
||||||
|
so cashu keeps using genuine httpx types everywhere else.
|
||||||
|
"""
|
||||||
|
|
||||||
|
AsyncClient = _ProxiesCompatAsyncClient
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> Any:
|
||||||
|
return getattr(httpx, name)
|
||||||
|
|
||||||
|
|
||||||
|
def _httpx_accepts_proxies() -> bool:
|
||||||
|
import inspect
|
||||||
|
|
||||||
|
try:
|
||||||
|
return "proxies" in inspect.signature(httpx.AsyncClient.__init__).parameters
|
||||||
|
except (TypeError, ValueError): # pragma: no cover - defensive
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def install_cashu_httpx_shim() -> bool:
|
||||||
|
"""Patch cashu's mint client to survive httpx>=0.28.
|
||||||
|
|
||||||
|
Returns True when the shim was installed, False when it wasn't needed.
|
||||||
|
Safe to call repeatedly.
|
||||||
|
"""
|
||||||
|
if _httpx_accepts_proxies():
|
||||||
|
return False
|
||||||
|
|
||||||
|
from cashu.wallet import v1_api
|
||||||
|
|
||||||
|
if isinstance(getattr(v1_api, "httpx", None), _HttpxNamespace):
|
||||||
|
return True
|
||||||
|
|
||||||
|
v1_api.httpx = _HttpxNamespace() # type: ignore[assignment]
|
||||||
|
return True
|
||||||
+255
-89
@@ -1,4 +1,3 @@
|
|||||||
import asyncio
|
|
||||||
import json
|
import json
|
||||||
import re
|
import re
|
||||||
import secrets
|
import secrets
|
||||||
@@ -6,12 +5,17 @@ from datetime import datetime, timezone
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from pydantic import BaseModel, RootModel
|
from pydantic import BaseModel, RootModel, field_validator
|
||||||
from pydantic.v1 import ValidationError as PydanticValidationError
|
from pydantic.v1 import ValidationError as PydanticValidationError
|
||||||
from sqlmodel import select
|
from sqlmodel import select
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from ..payment.models import _row_to_model, list_models
|
from ..payment.models import (
|
||||||
|
REQUIRED_PRICING_FIELDS,
|
||||||
|
_row_to_model,
|
||||||
|
list_models,
|
||||||
|
)
|
||||||
|
from ..payment.rates import BILLABLE_PRICING_FIELDS, coerce_rate
|
||||||
from ..proxy import refresh_model_maps, reinitialize_upstreams
|
from ..proxy import refresh_model_maps, reinitialize_upstreams
|
||||||
from ..wallet import fetch_all_balances, send_token, token_mint_url
|
from ..wallet import fetch_all_balances, send_token, token_mint_url
|
||||||
from . import vault
|
from . import vault
|
||||||
@@ -30,6 +34,7 @@ from .db import (
|
|||||||
from .db import (
|
from .db import (
|
||||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||||
)
|
)
|
||||||
|
from .exceptions import json_compliant
|
||||||
from .log_manager import log_manager
|
from .log_manager import log_manager
|
||||||
from .logging import get_logger
|
from .logging import get_logger
|
||||||
from .provider_slugs import allocate_unique_provider_slug
|
from .provider_slugs import allocate_unique_provider_slug
|
||||||
@@ -69,7 +74,9 @@ async def require_admin_api(request: Request) -> None:
|
|||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
result = await session.exec(select(CliToken).where(CliToken.token == token))
|
result = await session.exec(select(CliToken).where(CliToken.token == token))
|
||||||
cli_token = result.first()
|
cli_token = result.first()
|
||||||
if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts):
|
if cli_token and (
|
||||||
|
cli_token.expires_at is None or cli_token.expires_at > now_ts
|
||||||
|
):
|
||||||
cli_token.last_used_at = now_ts
|
cli_token.last_used_at = now_ts
|
||||||
session.add(cli_token)
|
session.add(cli_token)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
@@ -104,25 +111,28 @@ async def get_temporary_balances_api(
|
|||||||
)
|
)
|
||||||
total = count_result.one()
|
total = count_result.one()
|
||||||
|
|
||||||
# Aggregate totals across the whole (search-filtered) set, not just the
|
# Aggregate totals across the whole search-filtered set, not just this page.
|
||||||
# current page. Balance counts only parent (non-child) keys to avoid
|
balance_totals_result = await session.exec(
|
||||||
# double-counting, since child keys draw from their parent's balance.
|
|
||||||
totals_result = await session.exec(
|
|
||||||
select(
|
select(
|
||||||
|
func.coalesce(func.sum(ApiKey.balance), 0),
|
||||||
|
func.coalesce(func.sum(ApiKey.reserved_balance), 0),
|
||||||
func.coalesce(
|
func.coalesce(
|
||||||
func.sum(
|
func.sum(col(ApiKey.balance) - col(ApiKey.reserved_balance)), 0
|
||||||
case(
|
),
|
||||||
(col(ApiKey.parent_key_hash).is_(None), ApiKey.balance),
|
).where(*filters)
|
||||||
else_=0,
|
|
||||||
)
|
)
|
||||||
),
|
(
|
||||||
0,
|
total_balance,
|
||||||
),
|
total_reserved_balance,
|
||||||
|
total_available_balance,
|
||||||
|
) = balance_totals_result.one()
|
||||||
|
usage_totals_result = await session.exec(
|
||||||
|
select(
|
||||||
func.coalesce(func.sum(ApiKey.total_spent), 0),
|
func.coalesce(func.sum(ApiKey.total_spent), 0),
|
||||||
func.coalesce(func.sum(ApiKey.total_requests), 0),
|
func.coalesce(func.sum(ApiKey.total_requests), 0),
|
||||||
).where(*filters)
|
).where(*filters)
|
||||||
)
|
)
|
||||||
total_balance, total_spent, total_requests = totals_result.one()
|
total_spent, total_requests = usage_totals_result.one()
|
||||||
|
|
||||||
# Latest created first; keys with no created_at (legacy rows) sort last.
|
# Latest created first; keys with no created_at (legacy rows) sort last.
|
||||||
# Use an explicit CASE rather than relying on dialect NULL-ordering so
|
# Use an explicit CASE rather than relying on dialect NULL-ordering so
|
||||||
@@ -143,13 +153,12 @@ async def get_temporary_balances_api(
|
|||||||
{
|
{
|
||||||
"hashed_key": key.hashed_key,
|
"hashed_key": key.hashed_key,
|
||||||
"balance": key.balance,
|
"balance": key.balance,
|
||||||
|
"reserved_balance": key.reserved_balance,
|
||||||
|
"available_balance": key.total_balance,
|
||||||
"total_spent": key.total_spent,
|
"total_spent": key.total_spent,
|
||||||
"total_requests": key.total_requests,
|
"total_requests": key.total_requests,
|
||||||
"refund_address": key.refund_address,
|
"refund_address": key.refund_address,
|
||||||
"key_expiry_time": key.key_expiry_time,
|
"key_expiry_time": key.key_expiry_time,
|
||||||
"parent_key_hash": key.parent_key_hash,
|
|
||||||
"balance_limit": key.balance_limit,
|
|
||||||
"balance_limit_reset": key.balance_limit_reset,
|
|
||||||
"validity_date": key.validity_date,
|
"validity_date": key.validity_date,
|
||||||
"created_at": key.created_at,
|
"created_at": key.created_at,
|
||||||
}
|
}
|
||||||
@@ -158,6 +167,8 @@ async def get_temporary_balances_api(
|
|||||||
"total": total,
|
"total": total,
|
||||||
"totals": {
|
"totals": {
|
||||||
"total_balance": total_balance,
|
"total_balance": total_balance,
|
||||||
|
"total_reserved_balance": total_reserved_balance,
|
||||||
|
"total_available_balance": total_available_balance,
|
||||||
"total_spent": total_spent,
|
"total_spent": total_spent,
|
||||||
"total_requests": total_requests,
|
"total_requests": total_requests,
|
||||||
},
|
},
|
||||||
@@ -165,8 +176,6 @@ async def get_temporary_balances_api(
|
|||||||
|
|
||||||
|
|
||||||
class ApiKeyUpdate(BaseModel):
|
class ApiKeyUpdate(BaseModel):
|
||||||
balance_limit: int | None = None
|
|
||||||
balance_limit_reset: str | None = None
|
|
||||||
validity_date: int | None = None
|
validity_date: int | None = None
|
||||||
|
|
||||||
|
|
||||||
@@ -181,10 +190,6 @@ async def update_apikey(
|
|||||||
if not key:
|
if not key:
|
||||||
raise HTTPException(status_code=404, detail="API key not found")
|
raise HTTPException(status_code=404, detail="API key not found")
|
||||||
|
|
||||||
if update.balance_limit is not None:
|
|
||||||
key.balance_limit = update.balance_limit
|
|
||||||
if update.balance_limit_reset is not None:
|
|
||||||
key.balance_limit_reset = update.balance_limit_reset
|
|
||||||
if update.validity_date is not None:
|
if update.validity_date is not None:
|
||||||
key.validity_date = update.validity_date
|
key.validity_date = update.validity_date
|
||||||
|
|
||||||
@@ -194,8 +199,6 @@ async def update_apikey(
|
|||||||
|
|
||||||
return {
|
return {
|
||||||
"hashed_key": key.hashed_key,
|
"hashed_key": key.hashed_key,
|
||||||
"balance_limit": key.balance_limit,
|
|
||||||
"balance_limit_reset": key.balance_limit_reset,
|
|
||||||
"validity_date": key.validity_date,
|
"validity_date": key.validity_date,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -256,16 +259,12 @@ async def update_password(request: Request, password_update: PasswordUpdate) ->
|
|||||||
secret = await get_secret(session)
|
secret = await get_secret(session)
|
||||||
|
|
||||||
if not secret.admin_password_hash:
|
if not secret.admin_password_hash:
|
||||||
raise HTTPException(
|
raise HTTPException(status_code=500, detail="Admin password not configured")
|
||||||
status_code=500, detail="Admin password not configured"
|
|
||||||
)
|
|
||||||
|
|
||||||
if not vault.verify_password(
|
if not vault.verify_password(
|
||||||
password_update.current_password, secret.admin_password_hash
|
password_update.current_password, secret.admin_password_hash
|
||||||
):
|
):
|
||||||
raise HTTPException(
|
raise HTTPException(status_code=401, detail="Current password is incorrect")
|
||||||
status_code=401, detail="Current password is incorrect"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Validate new password
|
# Validate new password
|
||||||
new_password = password_update.new_password.strip()
|
new_password = password_update.new_password.strip()
|
||||||
@@ -483,6 +482,38 @@ class ModelCreate(BaseModel):
|
|||||||
enabled: bool = True
|
enabled: bool = True
|
||||||
forwarded_model_id: str | None = None
|
forwarded_model_id: str | None = None
|
||||||
|
|
||||||
|
@field_validator("pricing")
|
||||||
|
@classmethod
|
||||||
|
def _validate_pricing(cls, value: dict[str, object]) -> dict[str, object]:
|
||||||
|
"""Reject a rate that is malformed, non-finite, negative or not there.
|
||||||
|
|
||||||
|
A present-but-invalid rate would otherwise slip through: a non-numeric
|
||||||
|
string coerces to $0 on the read path (an unpriced-looking row), while a
|
||||||
|
negative or ``NaN``/``inf`` value is truthy and reads back as a real
|
||||||
|
price, so the model could be enabled and bill a nonsensical amount.
|
||||||
|
Surfacing a 422 reports the client bug as a client bug instead of
|
||||||
|
persisting it. Numeric strings (``"0.000005"``) stay valid, and so does
|
||||||
|
an omitted auxiliary rate — the stored JSON accepts both.
|
||||||
|
"""
|
||||||
|
for field in BILLABLE_PRICING_FIELDS:
|
||||||
|
if field not in value:
|
||||||
|
# ``dict.get`` cannot tell this from an explicit ``null``, so
|
||||||
|
# both were skipped and a row that ``Pricing`` cannot parse was
|
||||||
|
# written — and then raised out of the response that reads it
|
||||||
|
# back, after the row had been committed.
|
||||||
|
if field in REQUIRED_PRICING_FIELDS:
|
||||||
|
raise ValueError(f"{field} is required")
|
||||||
|
continue
|
||||||
|
# The shared coercion also absorbs the OverflowError an oversized
|
||||||
|
# integer raises, which pydantic does not convert into a validation
|
||||||
|
# error — unhandled it escaped as a 500 for a bad client value.
|
||||||
|
if coerce_rate(value[field]) is None:
|
||||||
|
raise ValueError(
|
||||||
|
f"{field} must be a finite, non-negative number, "
|
||||||
|
f"got {value[field]!r}"
|
||||||
|
)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
def _normalize_forwarded_model_id(value: str | None) -> str | None:
|
def _normalize_forwarded_model_id(value: str | None) -> str | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
@@ -609,9 +640,13 @@ async def get_provider_model(provider_id: str, model_id: str) -> dict[str, objec
|
|||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=404, detail="Model not found for this provider"
|
status_code=404, detail="Model not found for this provider"
|
||||||
)
|
)
|
||||||
return _row_to_model(
|
# Same duty as the listing this view is opened from: a stored rate that
|
||||||
|
# is not a usable number must be shown as it is, not encoded as `null`.
|
||||||
|
return json_compliant( # type: ignore[return-value]
|
||||||
|
_row_to_model(
|
||||||
row, apply_provider_fee=False, provider_fee=provider.provider_fee
|
row, apply_provider_fee=False, provider_fee=provider.provider_fee
|
||||||
).dict() # type: ignore
|
).dict()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@admin_router.delete(
|
@admin_router.delete(
|
||||||
@@ -870,29 +905,43 @@ class UpstreamProviderUpdateBySlug(BaseModel):
|
|||||||
provider_settings: dict | None = None
|
provider_settings: dict | None = None
|
||||||
|
|
||||||
|
|
||||||
async def _active_ppq_claim_in_session(session: AsyncSession, provider_id: int) -> bool:
|
async def _active_auto_topup_claim_in_session(
|
||||||
|
session: AsyncSession, provider_id: int, provider_type: str
|
||||||
|
) -> bool:
|
||||||
"""Check for an active claim inside the caller's transaction.
|
"""Check for an active claim inside the caller's transaction.
|
||||||
|
|
||||||
Must share the transaction of whatever destructive write it is guarding —
|
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
|
a check in its own session leaves a window for a worker to create the
|
||||||
claim between the check and the commit.
|
claim between the check and the commit.
|
||||||
"""
|
"""
|
||||||
from ..upstream.auto_topup import _ppq_state_id_for_provider
|
from ..upstream.auto_topup import (
|
||||||
|
_ppq_state_id_for_provider,
|
||||||
|
_routstr_state_id_for_provider,
|
||||||
|
)
|
||||||
|
|
||||||
claim = await session.get(CashuTransaction, _ppq_state_id_for_provider(provider_id))
|
state_id = (
|
||||||
|
_ppq_state_id_for_provider(provider_id)
|
||||||
|
if provider_type == "ppqai"
|
||||||
|
else _routstr_state_id_for_provider(provider_id)
|
||||||
|
)
|
||||||
|
claim = await session.get(CashuTransaction, state_id)
|
||||||
return claim is not None and not claim.collected and not claim.swept
|
return claim is not None and not claim.collected and not claim.swept
|
||||||
|
|
||||||
|
|
||||||
def _require_valid_ppq_auto_topup(
|
def _require_valid_auto_topup(provider_type: str, settings: dict | None) -> None:
|
||||||
provider_type: str, settings: dict | None
|
"""Reject auto top-up settings the worker would later refuse."""
|
||||||
) -> None:
|
from ..upstream.auto_topup import (
|
||||||
"""Reject PPQ auto top-up settings the worker would later refuse."""
|
validate_ppq_auto_topup_settings,
|
||||||
if provider_type != "ppqai":
|
validate_routstr_auto_topup_settings,
|
||||||
|
)
|
||||||
|
|
||||||
|
if provider_type == "ppqai":
|
||||||
|
problem = validate_ppq_auto_topup_settings(settings)
|
||||||
|
elif provider_type == "routstr":
|
||||||
|
problem = validate_routstr_auto_topup_settings(settings)
|
||||||
|
else:
|
||||||
return
|
return
|
||||||
|
|
||||||
from ..upstream.auto_topup import validate_ppq_auto_topup_settings
|
|
||||||
|
|
||||||
problem = validate_ppq_auto_topup_settings(settings)
|
|
||||||
if problem is not None:
|
if problem is not None:
|
||||||
raise HTTPException(status_code=400, detail=problem)
|
raise HTTPException(status_code=400, detail=problem)
|
||||||
|
|
||||||
@@ -917,16 +966,18 @@ async def _apply_provider_update(
|
|||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
provider_type_changed
|
provider_type_changed
|
||||||
and provider.provider_type == "ppqai"
|
and provider.provider_type in ("ppqai", "routstr")
|
||||||
and provider.id is not None
|
and provider.id is not None
|
||||||
and await _active_ppq_claim_in_session(session, provider.id)
|
and await _active_auto_topup_claim_in_session(
|
||||||
|
session, provider.id, provider.provider_type
|
||||||
|
)
|
||||||
):
|
):
|
||||||
# Changing the type would orphan the claim: the PPQ endpoints refuse
|
# Changing the type would orphan the claim: the claim endpoints refuse
|
||||||
# non-ppqai providers, so nobody could ever inspect or release it.
|
# providers of the wrong type, so nobody could inspect or release it.
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=409,
|
status_code=409,
|
||||||
detail=(
|
detail=(
|
||||||
"This provider has an active PPQ auto top-up claim. Release "
|
"This provider has an active auto top-up claim. Release "
|
||||||
"it before changing the provider type"
|
"it before changing the provider type"
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -977,7 +1028,7 @@ async def _apply_provider_update(
|
|||||||
except (json.JSONDecodeError, TypeError):
|
except (json.JSONDecodeError, TypeError):
|
||||||
effective_settings = None
|
effective_settings = None
|
||||||
if effective_settings is not None:
|
if effective_settings is not None:
|
||||||
_require_valid_ppq_auto_topup(provider.provider_type, effective_settings)
|
_require_valid_auto_topup(provider.provider_type, effective_settings)
|
||||||
if payload.provider_settings is not None:
|
if payload.provider_settings is not None:
|
||||||
provider.provider_settings = json.dumps(payload.provider_settings)
|
provider.provider_settings = json.dumps(payload.provider_settings)
|
||||||
|
|
||||||
@@ -1017,9 +1068,7 @@ async def create_upstream_provider(
|
|||||||
else:
|
else:
|
||||||
slug = await allocate_unique_provider_slug(session, payload.provider_type)
|
slug = await allocate_unique_provider_slug(session, payload.provider_type)
|
||||||
|
|
||||||
_require_valid_ppq_auto_topup(
|
_require_valid_auto_topup(payload.provider_type, payload.provider_settings)
|
||||||
payload.provider_type, payload.provider_settings
|
|
||||||
)
|
|
||||||
|
|
||||||
provider = UpstreamProviderRow(
|
provider = UpstreamProviderRow(
|
||||||
slug=slug,
|
slug=slug,
|
||||||
@@ -1078,9 +1127,7 @@ async def update_upstream_provider_by_slug(
|
|||||||
lookup = _validate_slug(payload.slug)
|
lookup = _validate_slug(payload.slug)
|
||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
result = await session.exec(
|
result = await session.exec(
|
||||||
select(UpstreamProviderRow).where(
|
select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup)
|
||||||
UpstreamProviderRow.slug == lookup
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
provider = result.first()
|
provider = result.first()
|
||||||
if not provider:
|
if not provider:
|
||||||
@@ -1117,16 +1164,19 @@ async def delete_upstream_provider(provider_id: str) -> dict[str, object]:
|
|||||||
# re-reads the provider inside its own transaction, so these two
|
# re-reads the provider inside its own transaction, so these two
|
||||||
# writes serialise — either the claim lands first and this 409s, or
|
# writes serialise — either the claim lands first and this 409s, or
|
||||||
# the delete lands first and the worker refuses to claim.
|
# the delete lands first and the worker refuses to claim.
|
||||||
if provider.provider_type == "ppqai" and await _active_ppq_claim_in_session(
|
if provider.provider_type in (
|
||||||
session, deleted_id
|
"ppqai",
|
||||||
|
"routstr",
|
||||||
|
) and await _active_auto_topup_claim_in_session(
|
||||||
|
session, deleted_id, provider.provider_type
|
||||||
):
|
):
|
||||||
# Deleting now would orphan the claim and any funds it tracks:
|
# Deleting now would orphan the claim and any funds it tracks:
|
||||||
# the PPQ endpoints 404 without the provider row, so the claim
|
# the claim endpoints 404 without the provider row, so the claim
|
||||||
# could never again be inspected or released.
|
# could never again be inspected or released.
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=409,
|
status_code=409,
|
||||||
detail=(
|
detail=(
|
||||||
"This provider has an active PPQ auto top-up claim. "
|
"This provider has an active auto top-up claim. "
|
||||||
"Resolve and release it before deleting the provider"
|
"Resolve and release it before deleting the provider"
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -1186,8 +1236,13 @@ async def get_provider_models(provider_id: str) -> dict[str, object]:
|
|||||||
"provider_type": provider.provider_type,
|
"provider_type": provider.provider_type,
|
||||||
"base_url": provider.base_url,
|
"base_url": provider.base_url,
|
||||||
},
|
},
|
||||||
"db_models": [m.dict() for m in db_models],
|
# This listing includes disabled models, so it is the one view that
|
||||||
"remote_models": [m.dict() for m in filtered_remote_models],
|
# still carries a row the served-catalog backstop holds back —
|
||||||
|
# including one whose stored rate is not a usable number. The
|
||||||
|
# encoder would report that rate as `null`, indistinguishable from a
|
||||||
|
# missing one; show the operator the value that needs fixing.
|
||||||
|
"db_models": [json_compliant(m.dict()) for m in db_models],
|
||||||
|
"remote_models": [json_compliant(m.dict()) for m in filtered_remote_models],
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -1316,12 +1371,7 @@ async def initiate_provider_topup(
|
|||||||
else {}
|
else {}
|
||||||
)
|
)
|
||||||
|
|
||||||
last_status_code = 500
|
# Quote creation is unsafe to retry without idempotency.
|
||||||
last_error_detail: object = "Failed to create top-up invoice"
|
|
||||||
|
|
||||||
# Some upstream Routstr nodes fail the first invoice request after warm-up
|
|
||||||
# and succeed immediately on retry. Retry once here so the UI stays single-click.
|
|
||||||
for attempt in range(2):
|
|
||||||
resp = await client.post(
|
resp = await client.post(
|
||||||
f"{clean_url}/v1/balance/lightning/invoice",
|
f"{clean_url}/v1/balance/lightning/invoice",
|
||||||
json=request_json,
|
json=request_json,
|
||||||
@@ -1343,23 +1393,15 @@ async def initiate_provider_topup(
|
|||||||
f"Upstream topup request failed: {resp.text}",
|
f"Upstream topup request failed: {resp.text}",
|
||||||
extra={
|
extra={
|
||||||
"provider_id": provider_id,
|
"provider_id": provider_id,
|
||||||
"attempt": attempt + 1,
|
|
||||||
"status_code": resp.status_code,
|
"status_code": resp.status_code,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
last_error_detail = resp.json()
|
error_detail: object = resp.json()
|
||||||
except Exception:
|
except Exception:
|
||||||
last_error_detail = resp.text
|
error_detail = resp.text
|
||||||
last_status_code = resp.status_code
|
|
||||||
|
|
||||||
if resp.status_code < 500 or attempt == 1:
|
|
||||||
break
|
|
||||||
|
|
||||||
await asyncio.sleep(0.2)
|
|
||||||
|
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=last_status_code, detail=last_error_detail
|
status_code=resp.status_code, detail=error_detail
|
||||||
)
|
)
|
||||||
|
|
||||||
upstream_instance = _instantiate_provider(provider)
|
upstream_instance = _instantiate_provider(provider)
|
||||||
@@ -1718,6 +1760,34 @@ async def get_logs_api(
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@admin_router.get(
|
||||||
|
"/api/logs/request/{request_id}", dependencies=[Depends(require_admin_api)]
|
||||||
|
)
|
||||||
|
async def get_logs_by_request_id_api(
|
||||||
|
request: Request,
|
||||||
|
request_id: str,
|
||||||
|
date: str | None = None,
|
||||||
|
limit: int = Query(default=200, ge=1, le=1000),
|
||||||
|
) -> dict[str, object]:
|
||||||
|
"""
|
||||||
|
Get every log entry belonging to a single request ID, oldest first.
|
||||||
|
"""
|
||||||
|
log_entries = log_manager.search_logs(
|
||||||
|
date=date,
|
||||||
|
request_id=request_id,
|
||||||
|
limit=limit,
|
||||||
|
)
|
||||||
|
log_entries.sort(key=lambda entry: str(entry.get("asctime", "")))
|
||||||
|
|
||||||
|
return {
|
||||||
|
"logs": log_entries,
|
||||||
|
"total": len(log_entries),
|
||||||
|
"request_id": request_id,
|
||||||
|
"date": date,
|
||||||
|
"limit": limit,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@admin_router.get("/api/logs/dates", dependencies=[Depends(require_admin_api)])
|
@admin_router.get("/api/logs/dates", dependencies=[Depends(require_admin_api)])
|
||||||
async def get_log_dates_api(request: Request) -> dict[str, object]:
|
async def get_log_dates_api(request: Request) -> dict[str, object]:
|
||||||
logs_dir = Path("logs")
|
logs_dir = Path("logs")
|
||||||
@@ -1811,6 +1881,87 @@ async def release_ppq_auto_topup_api(
|
|||||||
return {"ok": True, "released": True}
|
return {"ok": True, "released": True}
|
||||||
|
|
||||||
|
|
||||||
|
_ROUTSTR_RELEASE_ERRORS = {
|
||||||
|
"no_active_claim": "No active Routstr auto top-up claim to release",
|
||||||
|
"stale_state": "The claim changed since it was reviewed; reload and check again",
|
||||||
|
"claim_changed": (
|
||||||
|
"The claim changed while the release was being applied; reload and check again"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class ReleaseRoutstrAutoTopupRequest(BaseModel):
|
||||||
|
confirmed_peer_reconciled: bool
|
||||||
|
state_token: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
async def _require_routstr_provider(provider_id: int) -> UpstreamProviderRow:
|
||||||
|
async with create_session() as session:
|
||||||
|
provider = await session.get(UpstreamProviderRow, provider_id)
|
||||||
|
if provider is None:
|
||||||
|
raise HTTPException(status_code=404, detail="Provider not found")
|
||||||
|
if provider.provider_type != "routstr":
|
||||||
|
raise HTTPException(status_code=400, detail="Provider is not a Routstr node")
|
||||||
|
return provider
|
||||||
|
|
||||||
|
|
||||||
|
@admin_router.get(
|
||||||
|
"/api/upstream-providers/{provider_id}/routstr-auto-topup",
|
||||||
|
dependencies=[Depends(require_admin_api)],
|
||||||
|
)
|
||||||
|
async def get_routstr_auto_topup_api(provider_id: int) -> dict[str, object]:
|
||||||
|
await _require_routstr_provider(provider_id)
|
||||||
|
from ..upstream.auto_topup import get_routstr_auto_topup_state
|
||||||
|
|
||||||
|
return {"ok": True, **await get_routstr_auto_topup_state(provider_id)}
|
||||||
|
|
||||||
|
|
||||||
|
@admin_router.post(
|
||||||
|
"/api/upstream-providers/{provider_id}/routstr-auto-topup/release",
|
||||||
|
dependencies=[Depends(require_admin_api)],
|
||||||
|
)
|
||||||
|
async def release_routstr_auto_topup_api(
|
||||||
|
provider_id: int, payload: ReleaseRoutstrAutoTopupRequest
|
||||||
|
) -> dict[str, object]:
|
||||||
|
await _require_routstr_provider(provider_id)
|
||||||
|
if not payload.confirmed_peer_reconciled:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="Confirm the peer credited or returned the token before releasing",
|
||||||
|
)
|
||||||
|
|
||||||
|
from ..upstream.auto_topup import release_routstr_auto_topup_state
|
||||||
|
|
||||||
|
outcome = await release_routstr_auto_topup_state(
|
||||||
|
provider_id, state_token=payload.state_token
|
||||||
|
)
|
||||||
|
if not outcome.released:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=409, detail=_ROUTSTR_RELEASE_ERRORS[outcome.reason]
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.warning(
|
||||||
|
"Admin released Routstr auto top-up claim after manual reconciliation",
|
||||||
|
extra={"provider_id": provider_id, "state_token": payload.state_token},
|
||||||
|
)
|
||||||
|
return {"ok": True, "released": True}
|
||||||
|
|
||||||
|
|
||||||
|
def _transaction_status(tx: CashuTransaction) -> str:
|
||||||
|
"""An outgoing admin withdrawal ends at "issued": the node hands the bearer
|
||||||
|
token over and never learns whether it was redeemed, so its flags stay false
|
||||||
|
and it would otherwise read pending forever. Incoming admin rows are redeemed
|
||||||
|
by the node itself and keep the normal collected/swept lifecycle.
|
||||||
|
"""
|
||||||
|
if tx.swept:
|
||||||
|
return "swept"
|
||||||
|
if tx.collected:
|
||||||
|
return "collected"
|
||||||
|
if tx.source == "admin" and tx.type == "out":
|
||||||
|
return "issued"
|
||||||
|
return "pending"
|
||||||
|
|
||||||
|
|
||||||
@admin_router.get("/api/transactions", dependencies=[Depends(require_admin_api)])
|
@admin_router.get("/api/transactions", dependencies=[Depends(require_admin_api)])
|
||||||
async def get_transactions_api(
|
async def get_transactions_api(
|
||||||
type: str | None = None,
|
type: str | None = None,
|
||||||
@@ -1823,10 +1974,11 @@ async def get_transactions_api(
|
|||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
from sqlmodel import col, func
|
from sqlmodel import col, func
|
||||||
|
|
||||||
# Hide only the deterministic PPQ claim-lock rows. Append-only PPQ
|
# Hide only the deterministic claim-lock rows. Append-only PPQ payment
|
||||||
# payment rows remain visible as the audit trail for irreversible melts.
|
# rows and auto-topup token rows remain visible as the audit trail.
|
||||||
base = select(CashuTransaction).where(
|
base = select(CashuTransaction).where(
|
||||||
~col(CashuTransaction.id).like("ppq-auto-topup-%")
|
~col(CashuTransaction.id).like("ppq-auto-topup-%"),
|
||||||
|
~col(CashuTransaction.id).like("routstr-auto-topup-%"),
|
||||||
)
|
)
|
||||||
if type:
|
if type:
|
||||||
base = base.where(CashuTransaction.type == type)
|
base = base.where(CashuTransaction.type == type)
|
||||||
@@ -1840,13 +1992,25 @@ async def get_transactions_api(
|
|||||||
base = base.where(CashuTransaction.source == source)
|
base = base.where(CashuTransaction.source == source)
|
||||||
if status:
|
if status:
|
||||||
if status == "collected":
|
if status == "collected":
|
||||||
base = base.where(CashuTransaction.collected == True) # noqa: E712
|
base = base.where(
|
||||||
|
CashuTransaction.collected == True, # noqa: E712
|
||||||
|
CashuTransaction.swept == False, # noqa: E712
|
||||||
|
)
|
||||||
elif status == "swept":
|
elif status == "swept":
|
||||||
base = base.where(CashuTransaction.swept == True) # noqa: E712
|
base = base.where(CashuTransaction.swept == True) # noqa: E712
|
||||||
|
elif status == "issued":
|
||||||
|
base = base.where(
|
||||||
|
CashuTransaction.source == "admin",
|
||||||
|
CashuTransaction.type == "out",
|
||||||
|
CashuTransaction.collected == False, # noqa: E712
|
||||||
|
CashuTransaction.swept == False, # noqa: E712
|
||||||
|
)
|
||||||
elif status == "pending":
|
elif status == "pending":
|
||||||
base = base.where(
|
base = base.where(
|
||||||
CashuTransaction.collected == False, # noqa: E712
|
CashuTransaction.collected == False, # noqa: E712
|
||||||
CashuTransaction.swept == False, # noqa: E712
|
CashuTransaction.swept == False, # noqa: E712
|
||||||
|
(CashuTransaction.source != "admin")
|
||||||
|
| (CashuTransaction.type != "out"),
|
||||||
)
|
)
|
||||||
|
|
||||||
if search:
|
if search:
|
||||||
@@ -1873,15 +2037,17 @@ async def get_transactions_api(
|
|||||||
|
|
||||||
return {
|
return {
|
||||||
"transactions": [
|
"transactions": [
|
||||||
tx.dict(exclude={"sweep_started_at"}) for tx in transactions
|
{
|
||||||
|
**tx.dict(exclude={"sweep_started_at"}),
|
||||||
|
"status": _transaction_status(tx),
|
||||||
|
}
|
||||||
|
for tx in transactions
|
||||||
],
|
],
|
||||||
"total": total,
|
"total": total,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@admin_router.get(
|
@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)])
|
||||||
"/api/lightning-invoices", dependencies=[Depends(require_admin_api)]
|
|
||||||
)
|
|
||||||
async def get_lightning_invoices_api(
|
async def get_lightning_invoices_api(
|
||||||
status: str | None = None,
|
status: str | None = None,
|
||||||
purpose: str | None = None,
|
purpose: str | None = None,
|
||||||
|
|||||||
+179
-74
@@ -12,12 +12,11 @@ from typing import AsyncGenerator
|
|||||||
from alembic import command
|
from alembic import command
|
||||||
from alembic.config import Config
|
from alembic.config import Config
|
||||||
from alembic.util.exc import CommandError
|
from alembic.util.exc import CommandError
|
||||||
from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_
|
from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_, text
|
||||||
from sqlalchemy.engine import make_url
|
from sqlalchemy.engine import make_url
|
||||||
from sqlalchemy.exc import IntegrityError, OperationalError
|
from sqlalchemy.exc import IntegrityError, OperationalError
|
||||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||||
from sqlalchemy.ext.asyncio.engine import create_async_engine
|
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 import Field, Relationship, SQLModel, col, func, select, update
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
@@ -37,6 +36,9 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
|
|||||||
is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"}
|
is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"}
|
||||||
pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite
|
pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite
|
||||||
options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping}
|
options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping}
|
||||||
|
connect_args: dict[str, object] = {}
|
||||||
|
if is_sqlite and not is_memory_sqlite:
|
||||||
|
connect_args["timeout"] = settings.database_busy_timeout
|
||||||
if not is_memory_sqlite:
|
if not is_memory_sqlite:
|
||||||
options.update(
|
options.update(
|
||||||
pool_size=settings.database_pool_size,
|
pool_size=settings.database_pool_size,
|
||||||
@@ -51,9 +53,12 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
|
|||||||
"database_url_backend": backend,
|
"database_url_backend": backend,
|
||||||
"in_memory_sqlite": is_memory_sqlite,
|
"in_memory_sqlite": is_memory_sqlite,
|
||||||
**options,
|
**options,
|
||||||
|
"connect_args": connect_args,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
created_engine = create_async_engine(database_url, echo=False, **options)
|
created_engine = create_async_engine(
|
||||||
|
database_url, echo=False, connect_args=connect_args, **options
|
||||||
|
)
|
||||||
hold_warn_seconds = settings.database_pool_hold_warn_seconds
|
hold_warn_seconds = settings.database_pool_hold_warn_seconds
|
||||||
|
|
||||||
def record_pool_checkout(
|
def record_pool_checkout(
|
||||||
@@ -132,21 +137,6 @@ class ApiKey(SQLModel, table=True): # type: ignore
|
|||||||
default=None,
|
default=None,
|
||||||
description="Currency of the cashu-token",
|
description="Currency of the cashu-token",
|
||||||
)
|
)
|
||||||
parent_key_hash: str | None = Field(
|
|
||||||
default=None, foreign_key="api_keys.hashed_key", index=True
|
|
||||||
)
|
|
||||||
balance_limit: int | None = Field(
|
|
||||||
default=None,
|
|
||||||
description="Max spendable balance in msats for this key (mostly for child keys)",
|
|
||||||
)
|
|
||||||
balance_limit_reset: str | None = Field(
|
|
||||||
default=None,
|
|
||||||
description="Reset policy for balance limit (manual, daily, monthly, etc.)",
|
|
||||||
)
|
|
||||||
balance_limit_reset_date: int | None = Field(
|
|
||||||
default=None,
|
|
||||||
description="Unix timestamp of the last time the balance limit was reset",
|
|
||||||
)
|
|
||||||
validity_date: int | None = Field(
|
validity_date: int | None = Field(
|
||||||
default=None,
|
default=None,
|
||||||
description="Unix timestamp after which the key is no longer valid",
|
description="Unix timestamp after which the key is no longer valid",
|
||||||
@@ -171,6 +161,55 @@ async def reset_all_reserved_balances(session: AsyncSession) -> None:
|
|||||||
logger.info("Reset reserved balances on startup")
|
logger.info("Reset reserved balances on startup")
|
||||||
|
|
||||||
|
|
||||||
|
async def _transition_stale_reservation(
|
||||||
|
session: AsyncSession, reservation_id: str, cutoff: int
|
||||||
|
) -> bool:
|
||||||
|
"""Mark one reservation released iff its lease is still older than cutoff.
|
||||||
|
|
||||||
|
``created_at`` doubles as the heartbeat lease timestamp, so the guard must
|
||||||
|
be part of this update: a reservation renewed between the sweeper's select
|
||||||
|
and this transition is in flight and must survive.
|
||||||
|
"""
|
||||||
|
transition = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ReservationRelease)
|
||||||
|
.where(col(ReservationRelease.id) == reservation_id)
|
||||||
|
.where(col(ReservationRelease.status) == "active")
|
||||||
|
.where(col(ReservationRelease.created_at) < cutoff)
|
||||||
|
.values(status="released")
|
||||||
|
)
|
||||||
|
return bool(transition.rowcount == 1)
|
||||||
|
|
||||||
|
|
||||||
|
async def _release_legacy_aggregate(
|
||||||
|
session: AsyncSession,
|
||||||
|
key_hash: str,
|
||||||
|
observed_reserved: int,
|
||||||
|
observed_reserved_at: int | None,
|
||||||
|
) -> bool:
|
||||||
|
"""Zero one legacy aggregate reservation iff it is exactly as observed.
|
||||||
|
|
||||||
|
A new reservation committing between the sweeper's read and this update
|
||||||
|
changes ``reserved_balance``/``reserved_at`` in the same transaction that
|
||||||
|
creates its durable row, so this compare-and-swap fails instead of erasing
|
||||||
|
the newcomer's reserved funds.
|
||||||
|
"""
|
||||||
|
if observed_reserved <= 0:
|
||||||
|
return False
|
||||||
|
reserved_at_guard = (
|
||||||
|
col(ApiKey.reserved_at).is_(None)
|
||||||
|
if observed_reserved_at is None
|
||||||
|
else col(ApiKey.reserved_at) == observed_reserved_at
|
||||||
|
)
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.where(col(ApiKey.reserved_balance) == observed_reserved)
|
||||||
|
.where(reserved_at_guard)
|
||||||
|
.values(reserved_balance=0, reserved_at=None)
|
||||||
|
)
|
||||||
|
return bool(result.rowcount == 1)
|
||||||
|
|
||||||
|
|
||||||
async def release_stale_reservations(
|
async def release_stale_reservations(
|
||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
max_age_seconds: int,
|
max_age_seconds: int,
|
||||||
@@ -191,25 +230,22 @@ async def release_stale_reservations(
|
|||||||
col(ReservationRelease.billing_key_hash) == key_hash,
|
col(ReservationRelease.billing_key_hash) == key_hash,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
reservations = (await session.exec(query)).all()
|
# Capture primitives: a repair rollback below would expire ORM instances.
|
||||||
|
reservation_rows = [
|
||||||
|
(r.id, r.key_hash, r.billing_key_hash, r.reserved_msats)
|
||||||
|
for r in (await session.exec(query)).all()
|
||||||
|
]
|
||||||
released = 0
|
released = 0
|
||||||
|
|
||||||
for reservation in reservations:
|
for res_id, res_key_hash, res_billing_hash, res_msats in reservation_rows:
|
||||||
transition = await session.exec( # type: ignore[call-overload]
|
if not await _transition_stale_reservation(session, res_id, cutoff):
|
||||||
update(ReservationRelease)
|
|
||||||
.where(col(ReservationRelease.id) == reservation.id)
|
|
||||||
.where(col(ReservationRelease.status) == "active")
|
|
||||||
.values(status="released")
|
|
||||||
)
|
|
||||||
if transition.rowcount != 1:
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
values = {
|
values = {
|
||||||
"reserved_balance": col(ApiKey.reserved_balance)
|
"reserved_balance": col(ApiKey.reserved_balance) - res_msats,
|
||||||
- reservation.reserved_msats,
|
|
||||||
"reserved_at": case(
|
"reserved_at": case(
|
||||||
(
|
(
|
||||||
col(ApiKey.reserved_balance) - reservation.reserved_msats > 0,
|
col(ApiKey.reserved_balance) - res_msats > 0,
|
||||||
col(ApiKey.reserved_at),
|
col(ApiKey.reserved_at),
|
||||||
),
|
),
|
||||||
else_=None,
|
else_=None,
|
||||||
@@ -217,24 +253,42 @@ async def release_stale_reservations(
|
|||||||
}
|
}
|
||||||
parent_result = await session.exec( # type: ignore[call-overload]
|
parent_result = await session.exec( # type: ignore[call-overload]
|
||||||
update(ApiKey)
|
update(ApiKey)
|
||||||
.where(col(ApiKey.hashed_key) == reservation.billing_key_hash)
|
.where(col(ApiKey.hashed_key) == res_billing_hash)
|
||||||
.where(col(ApiKey.reserved_balance) >= reservation.reserved_msats)
|
.where(col(ApiKey.reserved_balance) >= res_msats)
|
||||||
.values(**values)
|
.values(**values)
|
||||||
)
|
)
|
||||||
if parent_result.rowcount != 1:
|
aggregates_ok = parent_result.rowcount == 1
|
||||||
await session.rollback()
|
if aggregates_ok and res_billing_hash != res_key_hash:
|
||||||
return 0
|
|
||||||
|
|
||||||
if reservation.billing_key_hash != reservation.key_hash:
|
|
||||||
child_result = await session.exec( # type: ignore[call-overload]
|
child_result = await session.exec( # type: ignore[call-overload]
|
||||||
update(ApiKey)
|
update(ApiKey)
|
||||||
.where(col(ApiKey.hashed_key) == reservation.key_hash)
|
.where(col(ApiKey.hashed_key) == res_key_hash)
|
||||||
.where(col(ApiKey.reserved_balance) >= reservation.reserved_msats)
|
.where(col(ApiKey.reserved_balance) >= res_msats)
|
||||||
.values(**values)
|
.values(**values)
|
||||||
)
|
)
|
||||||
if child_result.rowcount != 1:
|
aggregates_ok = child_result.rowcount == 1
|
||||||
|
|
||||||
|
if not aggregates_ok:
|
||||||
|
# The aggregates no longer hold this reservation's msats — the
|
||||||
|
# durable row is corrupt. Repair by terminalizing it WITHOUT
|
||||||
|
# subtracting uncertain aggregates (legacy cleanup below reconciles
|
||||||
|
# any stale remainder) and keep sweeping the rest of the batch:
|
||||||
|
# one corrupt row must not poison all stale cleanup.
|
||||||
await session.rollback()
|
await session.rollback()
|
||||||
return 0
|
if await _transition_stale_reservation(session, res_id, cutoff):
|
||||||
|
await session.commit()
|
||||||
|
released += 1
|
||||||
|
logger.error(
|
||||||
|
"Released corrupt stale reservation without aggregate subtraction",
|
||||||
|
extra={
|
||||||
|
"reservation_id": res_id,
|
||||||
|
"billing_key_hash": res_billing_hash[:8] + "...",
|
||||||
|
"reserved_msats": res_msats,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
# Commit each release on its own so a later corrupt record's rollback
|
||||||
|
# cannot discard the healthy releases already processed in this batch.
|
||||||
|
await session.commit()
|
||||||
released += 1
|
released += 1
|
||||||
|
|
||||||
# Rolling upgrades can leave aggregate reservations created before durable
|
# Rolling upgrades can leave aggregate reservations created before durable
|
||||||
@@ -246,16 +300,13 @@ async def release_stale_reservations(
|
|||||||
col(ApiKey.reserved_at) < cutoff
|
col(ApiKey.reserved_at) < cutoff
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
legacy_query = legacy_query.where(
|
legacy_query = legacy_query.where(col(ApiKey.hashed_key) == key_hash).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)
|
or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff)
|
||||||
)
|
)
|
||||||
|
|
||||||
for legacy_key in (await session.exec(legacy_query)).all():
|
for legacy_key in (await session.exec(legacy_query)).all():
|
||||||
|
observed_reserved = legacy_key.reserved_balance
|
||||||
|
observed_reserved_at = legacy_key.reserved_at
|
||||||
active_owner = (
|
active_owner = (
|
||||||
await session.exec(
|
await session.exec(
|
||||||
select(ReservationRelease.id)
|
select(ReservationRelease.id)
|
||||||
@@ -272,9 +323,9 @@ async def release_stale_reservations(
|
|||||||
).first()
|
).first()
|
||||||
if active_owner is not None:
|
if active_owner is not None:
|
||||||
continue
|
continue
|
||||||
legacy_key.reserved_balance = 0
|
if await _release_legacy_aggregate(
|
||||||
legacy_key.reserved_at = None
|
session, legacy_key.hashed_key, observed_reserved, observed_reserved_at
|
||||||
session.add(legacy_key)
|
):
|
||||||
released += 1
|
released += 1
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
@@ -290,21 +341,15 @@ async def release_stale_reservations(
|
|||||||
|
|
||||||
|
|
||||||
async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> int:
|
async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> int:
|
||||||
"""Delete dead parentless API keys; return the count removed.
|
"""Delete dead API keys; return the count removed.
|
||||||
|
|
||||||
Dead = 0 balance/reservation/spend/requests, older than the grace period,
|
Dead = 0 balance/reservation/spend/requests, older than the grace
|
||||||
no parent, no children, no invoice that could still settle. Cashu rows are
|
period, no invoice that could still settle. Cashu rows are
|
||||||
unlinked (not deleted) first to keep the audit trail.
|
unlinked (not deleted) first to keep the audit trail.
|
||||||
"""
|
"""
|
||||||
now = int(time.time())
|
now = int(time.time())
|
||||||
cutoff = now - min_age_seconds
|
cutoff = now - min_age_seconds
|
||||||
|
|
||||||
child = aliased(ApiKey)
|
|
||||||
has_children = (
|
|
||||||
select(child.hashed_key).where(
|
|
||||||
col(child.parent_key_hash) == col(ApiKey.hashed_key)
|
|
||||||
)
|
|
||||||
).exists()
|
|
||||||
# An expired invoice stays creditable for the grace window, and crediting it
|
# An expired invoice stays creditable for the grace window, and crediting it
|
||||||
# after its target key is gone strands the payment at the mint.
|
# after its target key is gone strands the payment at the mint.
|
||||||
settleable_invoice = (
|
settleable_invoice = (
|
||||||
@@ -322,16 +367,22 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
|
|||||||
)
|
)
|
||||||
).exists()
|
).exists()
|
||||||
|
|
||||||
|
has_refund_claim = (
|
||||||
|
select(Refund.id).where(
|
||||||
|
col(Refund.api_key_hashed_key) == col(ApiKey.hashed_key)
|
||||||
|
)
|
||||||
|
).exists()
|
||||||
|
|
||||||
eligible_hashes = (
|
eligible_hashes = (
|
||||||
select(ApiKey.hashed_key)
|
select(ApiKey.hashed_key)
|
||||||
.where(col(ApiKey.balance) == 0)
|
.where(col(ApiKey.balance) == 0)
|
||||||
.where(col(ApiKey.reserved_balance) == 0)
|
.where(col(ApiKey.reserved_balance) == 0)
|
||||||
.where(col(ApiKey.total_spent) == 0)
|
.where(col(ApiKey.total_spent) == 0)
|
||||||
.where(col(ApiKey.total_requests) == 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(~settleable_invoice)
|
.where(~settleable_invoice)
|
||||||
.where(~has_children)
|
# refunds holds a non-null FK to the key.
|
||||||
|
.where(~has_refund_claim)
|
||||||
)
|
)
|
||||||
|
|
||||||
# Unlink transactions rather than cascade-deleting them, so the financial
|
# Unlink transactions rather than cascade-deleting them, so the financial
|
||||||
@@ -386,9 +437,10 @@ class ModelRow(SQLModel, table=True): # type: ignore
|
|||||||
class ModelPathRow(SQLModel, table=True): # type: ignore
|
class ModelPathRow(SQLModel, table=True): # type: ignore
|
||||||
"""Upstream provider path a model is reachable through.
|
"""Upstream provider path a model is reachable through.
|
||||||
|
|
||||||
Discovery/visibility data only. ``model_id`` is intentionally NOT globally
|
Discovery data plus provider-specific model metadata. ``model_id`` is
|
||||||
unique: it is the client-visible ``/v1/models`` id (``forwarded_model_id or
|
intentionally NOT globally unique: it is the client-visible ``/v1/models``
|
||||||
id``) grouped across every provider that exposes the model. A single model
|
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
|
can therefore have several rows — one per direct provider path plus one per
|
||||||
OpenRouter sub-provider endpoint.
|
OpenRouter sub-provider endpoint.
|
||||||
"""
|
"""
|
||||||
@@ -425,6 +477,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore
|
|||||||
endpoint_name: str | None = Field(
|
endpoint_name: str | None = Field(
|
||||||
default=None, description="Human-readable endpoint display name"
|
default=None, description="Human-readable endpoint display name"
|
||||||
)
|
)
|
||||||
|
model_metadata: str = Field(
|
||||||
|
default="{}",
|
||||||
|
description="JSON model metadata specific to this provider path",
|
||||||
|
)
|
||||||
upstream_provider_id: int = Field(
|
upstream_provider_id: int = Field(
|
||||||
index=True,
|
index=True,
|
||||||
foreign_key="upstream_providers.id",
|
foreign_key="upstream_providers.id",
|
||||||
@@ -470,14 +526,6 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
|
|||||||
)
|
)
|
||||||
expires_at: int = Field(description="Unix timestamp when invoice expires")
|
expires_at: int = Field(description="Unix timestamp when invoice expires")
|
||||||
paid_at: int | None = Field(default=None, description="Unix timestamp when paid")
|
paid_at: int | None = Field(default=None, description="Unix timestamp when paid")
|
||||||
balance_limit: int | None = Field(
|
|
||||||
default=None,
|
|
||||||
description="Max spendable msats for the created key",
|
|
||||||
)
|
|
||||||
balance_limit_reset: str | None = Field(
|
|
||||||
default=None,
|
|
||||||
description="Reset policy for balance limit (daily, weekly, monthly)",
|
|
||||||
)
|
|
||||||
validity_date: int | None = Field(
|
validity_date: int | None = Field(
|
||||||
default=None,
|
default=None,
|
||||||
description="Unix timestamp after which the created key expires",
|
description="Unix timestamp after which the created key expires",
|
||||||
@@ -520,6 +568,53 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
REFUND_OPEN_STATUSES = ("pending", "ambiguous")
|
||||||
|
|
||||||
|
# Debited from the key but neither paid out nor restored, so still owed.
|
||||||
|
REFUND_UNRESOLVED_STATUSES = ("pending", "ambiguous", "stuck")
|
||||||
|
|
||||||
|
_REFUND_OPEN_PREDICATE = "status IN ('pending', 'ambiguous')"
|
||||||
|
|
||||||
|
|
||||||
|
class Refund(SQLModel, table=True): # type: ignore
|
||||||
|
"""One payout claim; the partial unique index allows one open claim per key."""
|
||||||
|
|
||||||
|
__tablename__ = "refunds"
|
||||||
|
__table_args__ = (
|
||||||
|
Index(
|
||||||
|
"ux_refunds_open_per_key",
|
||||||
|
"api_key_hashed_key",
|
||||||
|
unique=True,
|
||||||
|
sqlite_where=text(_REFUND_OPEN_PREDICATE),
|
||||||
|
postgresql_where=text(_REFUND_OPEN_PREDICATE),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex)
|
||||||
|
api_key_hashed_key: str = Field(foreign_key="api_keys.hashed_key", index=True)
|
||||||
|
method: str = Field(description="Payout method: lightning or cashu")
|
||||||
|
destination: str | None = Field(
|
||||||
|
default=None, description="Lightning address or LNURL, NULL for cashu"
|
||||||
|
)
|
||||||
|
amount_msats: int = Field(description="Balance debited when the claim opened")
|
||||||
|
unit: str = Field(description="Mint unit the payout is denominated in")
|
||||||
|
mint_url: str = Field(description="Mint the payout is drawn from")
|
||||||
|
status: str = Field(
|
||||||
|
default="pending",
|
||||||
|
index=True,
|
||||||
|
description="pending, paid, failed, ambiguous, or stuck",
|
||||||
|
)
|
||||||
|
quote_id: str | None = Field(
|
||||||
|
default=None, description="Melt quote id, for reconciling an ambiguous payout"
|
||||||
|
)
|
||||||
|
token: str | None = Field(default=None, description="Issued cashu token")
|
||||||
|
claimed_at: int | None = Field(
|
||||||
|
default=None, description="Reconciler lease timestamp"
|
||||||
|
)
|
||||||
|
created_at: int = Field(default_factory=lambda: int(time.time()))
|
||||||
|
updated_at: int = Field(default_factory=lambda: int(time.time()))
|
||||||
|
|
||||||
|
|
||||||
async def store_cashu_transaction(
|
async def store_cashu_transaction(
|
||||||
token: str,
|
token: str,
|
||||||
amount: int,
|
amount: int,
|
||||||
@@ -911,8 +1006,18 @@ async def complete_routstr_fee_payout(
|
|||||||
|
|
||||||
|
|
||||||
async def total_user_liability(db_session: AsyncSession) -> int:
|
async def total_user_liability(db_session: AsyncSession) -> int:
|
||||||
"""Return all outstanding API-key balances in millisatoshis."""
|
"""Return all outstanding user funds in millisatoshis.
|
||||||
result = await db_session.exec(select(func.sum(ApiKey.balance)))
|
|
||||||
|
Key balances and unresolved refunds are summed in one statement so a
|
||||||
|
claim opened between two reads cannot be missed by both.
|
||||||
|
"""
|
||||||
|
key_balances = select(func.coalesce(func.sum(ApiKey.balance), 0)).scalar_subquery()
|
||||||
|
unresolved_refunds = (
|
||||||
|
select(func.coalesce(func.sum(Refund.amount_msats), 0))
|
||||||
|
.where(col(Refund.status).in_(REFUND_UNRESOLVED_STATUSES))
|
||||||
|
.scalar_subquery()
|
||||||
|
)
|
||||||
|
result = await db_session.exec(select(key_balances + unresolved_refunds))
|
||||||
return int(result.one() or 0)
|
return int(result.one() or 0)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,8 @@
|
|||||||
|
import math
|
||||||
|
|
||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
|
from fastapi.encoders import jsonable_encoder
|
||||||
|
from fastapi.exceptions import RequestValidationError
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
|
|
||||||
from .logging import get_logger
|
from .logging import get_logger
|
||||||
@@ -30,6 +34,26 @@ class UpstreamError(Exception):
|
|||||||
super().__init__(message)
|
super().__init__(message)
|
||||||
|
|
||||||
|
|
||||||
|
class EhbpTimeoutError(UpstreamError):
|
||||||
|
"""Raised when an EHBP upstream times out waiting for a response.
|
||||||
|
|
||||||
|
Distinct from a generic :class:`UpstreamError` so callers can map the
|
||||||
|
failure to a ``504 Gateway Timeout`` with a stable ``UPSTREAM_TIMEOUT``
|
||||||
|
code instead of a misleading ``500`` internal server error.
|
||||||
|
|
||||||
|
``details`` carries optional structured, redaction-safe context and is
|
||||||
|
forwarded to the client by ``create_upstream_error_response``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, message: str, details: dict[str, object] | None = None):
|
||||||
|
super().__init__(
|
||||||
|
message,
|
||||||
|
status_code=504,
|
||||||
|
code="UPSTREAM_TIMEOUT",
|
||||||
|
details=details,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse:
|
async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse:
|
||||||
"""Handle HTTP exceptions and include request ID in response."""
|
"""Handle HTTP exceptions and include request ID in response."""
|
||||||
request_id = getattr(request.state, "request_id", "unknown")
|
request_id = getattr(request.state, "request_id", "unknown")
|
||||||
@@ -40,15 +64,25 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon
|
|||||||
path = request.url.path
|
path = request.url.path
|
||||||
|
|
||||||
# 4xx is client behaviour; the uvicorn access log already records it.
|
# 4xx is client behaviour; the uvicorn access log already records it.
|
||||||
# Only 5xx warrants a server-side warning/error log here.
|
|
||||||
if status_code >= 500:
|
if status_code >= 500:
|
||||||
logger.error(
|
error_type = None
|
||||||
|
if isinstance(detail, dict):
|
||||||
|
error = detail.get("error")
|
||||||
|
if isinstance(error, dict):
|
||||||
|
error_type = error.get("type")
|
||||||
|
log = (
|
||||||
|
logger.warning
|
||||||
|
if error_type in {"mint_unreachable", "mint_rate_limited"}
|
||||||
|
else logger.error
|
||||||
|
)
|
||||||
|
log(
|
||||||
f"HTTP {status_code} on {path}: {detail}",
|
f"HTTP {status_code} on {path}: {detail}",
|
||||||
extra={
|
extra={
|
||||||
"request_id": request_id,
|
"request_id": request_id,
|
||||||
"status_code": status_code,
|
"status_code": status_code,
|
||||||
"detail": detail,
|
"detail": detail,
|
||||||
"path": path,
|
"path": path,
|
||||||
|
"error_type": error_type,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -61,6 +95,45 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon
|
|||||||
return JSONResponse(status_code=status_code, content=content)
|
return JSONResponse(status_code=status_code, content=content)
|
||||||
|
|
||||||
|
|
||||||
|
def json_compliant(value: object) -> object:
|
||||||
|
"""Render non-finite floats as text so a reply carrying them can serialize.
|
||||||
|
|
||||||
|
``json`` parses the bare ``NaN``/``Infinity``/``-Infinity`` literals into
|
||||||
|
real floats, so a request body — and a stored row written from one — may
|
||||||
|
hold one anywhere. ``JSONResponse`` encodes with ``allow_nan=False`` and
|
||||||
|
raises on them, which would turn a reply that merely *quotes* the offending
|
||||||
|
value into a 500.
|
||||||
|
"""
|
||||||
|
if isinstance(value, float) and not math.isfinite(value):
|
||||||
|
return repr(value)
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return {key: json_compliant(item) for key, item in value.items()}
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return [json_compliant(item) for item in value]
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
async def validation_exception_handler(
|
||||||
|
request: Request, exc: Exception
|
||||||
|
) -> JSONResponse:
|
||||||
|
"""Answer a request-validation failure with a 422 that always serializes.
|
||||||
|
|
||||||
|
Pydantic echoes the rejected value back in each error's ``input`` field. A
|
||||||
|
non-finite float there breaks the encoder, so the 422 escapes as a 500 and
|
||||||
|
reports a client's bad rate as a server fault.
|
||||||
|
"""
|
||||||
|
request_id = getattr(request.state, "request_id", "unknown")
|
||||||
|
errors = exc.errors() if isinstance(exc, RequestValidationError) else []
|
||||||
|
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=422,
|
||||||
|
content={
|
||||||
|
"detail": json_compliant(jsonable_encoder(errors)),
|
||||||
|
"request_id": request_id,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def general_exception_handler(request: Request, exc: Exception) -> JSONResponse:
|
async def general_exception_handler(request: Request, exc: Exception) -> JSONResponse:
|
||||||
"""Handle general exceptions and include request ID in response."""
|
"""Handle general exceptions and include request ID in response."""
|
||||||
request_id = getattr(request.state, "request_id", "unknown")
|
request_id = getattr(request.state, "request_id", "unknown")
|
||||||
|
|||||||
+30
-12
@@ -16,7 +16,7 @@ DO NOT modify or remove these messages without updating the usage tracking logic
|
|||||||
- The 'token_cost', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics
|
- The 'token_cost', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics
|
||||||
|
|
||||||
3. "Max cost payment finalized" (INFO) - routstr/auth.py
|
3. "Max cost payment finalized" (INFO) - routstr/auth.py
|
||||||
- Used as the successful completion fallback when token usage is unavailable
|
- Used for explicit flat-price/MaxCostData settlements; missing usage alone must not create this charge
|
||||||
- The 'charged_amount', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics
|
- The 'charged_amount', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics
|
||||||
|
|
||||||
4. "Payment processed successfully" (INFO) - routstr/auth.py
|
4. "Payment processed successfully" (INFO) - routstr/auth.py
|
||||||
@@ -51,7 +51,7 @@ from pythonjsonlogger import jsonlogger
|
|||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
from rich.logging import RichHandler
|
from rich.logging import RichHandler
|
||||||
|
|
||||||
from .redaction import redact_obj, redact_org_ids
|
from .redaction import redact_field, redact_org_ids
|
||||||
|
|
||||||
# Only use RichHandler when stdout is a real TTY. In non-TTY contexts
|
# Only use RichHandler when stdout is a real TTY. In non-TTY contexts
|
||||||
# (docker logs, pipes, CI) Rich pads every line to width and wraps long
|
# (docker logs, pipes, CI) Rich pads every line to width and wraps long
|
||||||
@@ -100,8 +100,10 @@ class DailyRotatingFileHandler(logging.handlers.TimedRotatingFileHandler):
|
|||||||
self.baseFilename = new_filename
|
self.baseFilename = new_filename
|
||||||
self.current_date = new_date
|
self.current_date = new_date
|
||||||
|
|
||||||
# FIX ME: not sure if we need this
|
# `backupCount` alone never prunes these files: the base filename moves
|
||||||
# self._cleanup_old_files()
|
# with the date, so the inherited rollover finds no siblings to expire
|
||||||
|
# and every day of logged credentials is retained indefinitely.
|
||||||
|
self._cleanup_old_files()
|
||||||
|
|
||||||
if not self.delay:
|
if not self.delay:
|
||||||
self.stream = self._open()
|
self.stream = self._open()
|
||||||
@@ -182,6 +184,18 @@ class RequestIdFilter(logging.Filter):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
class ClientAppFilter(logging.Filter):
|
||||||
|
"""Attach request-local app attribution to log records."""
|
||||||
|
|
||||||
|
def filter(self, record: logging.LogRecord) -> bool:
|
||||||
|
# Import here to avoid circular imports
|
||||||
|
from .middleware import UNKNOWN_CLIENT_APP, client_app_context
|
||||||
|
|
||||||
|
client_app = client_app_context.get(None)
|
||||||
|
record.client_app = client_app if client_app else UNKNOWN_CLIENT_APP
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
# Standard ``LogRecord`` attributes that are never user-supplied ``extra``
|
# Standard ``LogRecord`` attributes that are never user-supplied ``extra``
|
||||||
# fields; skipped when redacting structured extras (``msg``/``message`` are
|
# fields; skipped when redacting structured extras (``msg``/``message`` are
|
||||||
# handled separately above).
|
# handled separately above).
|
||||||
@@ -260,13 +274,11 @@ class SecurityFilter(logging.Filter):
|
|||||||
|
|
||||||
# Structured `extra={...}` fields are emitted by the JSON formatter
|
# Structured `extra={...}` fields are emitted by the JSON formatter
|
||||||
# straight from the record dict and never pass through the message
|
# straight from the record dict and never pass through the message
|
||||||
# formatting above. Redact organization IDs from any string-valued
|
# formatting above, so they need their own recursive pass.
|
||||||
# extra so they cannot leak via structured logs.
|
|
||||||
for attr, value in list(record.__dict__.items()):
|
for attr, value in list(record.__dict__.items()):
|
||||||
if attr in _NON_EXTRA_RECORD_ATTRS:
|
if attr in _NON_EXTRA_RECORD_ATTRS:
|
||||||
continue
|
continue
|
||||||
if isinstance(value, (str, dict, list, tuple)):
|
record.__dict__[attr] = redact_field(attr, value)
|
||||||
record.__dict__[attr] = redact_obj(value)
|
|
||||||
|
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
@@ -323,7 +335,7 @@ def setup_logging() -> None:
|
|||||||
"rich_tracebacks": True,
|
"rich_tracebacks": True,
|
||||||
"markup": True,
|
"markup": True,
|
||||||
"console": _console,
|
"console": _console,
|
||||||
"filters": ["request_id_filter", "security_filter"],
|
"filters": ["request_id_filter", "client_app_filter", "security_filter"],
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
console_handler = {
|
console_handler = {
|
||||||
@@ -331,7 +343,7 @@ def setup_logging() -> None:
|
|||||||
"level": log_level,
|
"level": log_level,
|
||||||
"formatter": "plain",
|
"formatter": "plain",
|
||||||
"stream": "ext://sys.stdout",
|
"stream": "ext://sys.stdout",
|
||||||
"filters": ["request_id_filter", "security_filter"],
|
"filters": ["request_id_filter", "client_app_filter", "security_filter"],
|
||||||
}
|
}
|
||||||
|
|
||||||
LOGGING_CONFIG = {
|
LOGGING_CONFIG = {
|
||||||
@@ -340,7 +352,7 @@ def setup_logging() -> None:
|
|||||||
"formatters": {
|
"formatters": {
|
||||||
"json": {
|
"json": {
|
||||||
"()": jsonlogger.JsonFormatter,
|
"()": jsonlogger.JsonFormatter,
|
||||||
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s",
|
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s %(client_app)s",
|
||||||
"datefmt": "%Y-%m-%d %H:%M:%S",
|
"datefmt": "%Y-%m-%d %H:%M:%S",
|
||||||
},
|
},
|
||||||
"plain": {
|
"plain": {
|
||||||
@@ -351,6 +363,7 @@ def setup_logging() -> None:
|
|||||||
"filters": {
|
"filters": {
|
||||||
"version_filter": {"()": VersionFilter},
|
"version_filter": {"()": VersionFilter},
|
||||||
"request_id_filter": {"()": RequestIdFilter},
|
"request_id_filter": {"()": RequestIdFilter},
|
||||||
|
"client_app_filter": {"()": ClientAppFilter},
|
||||||
"security_filter": {"()": SecurityFilter},
|
"security_filter": {"()": SecurityFilter},
|
||||||
},
|
},
|
||||||
"handlers": {
|
"handlers": {
|
||||||
@@ -364,7 +377,12 @@ def setup_logging() -> None:
|
|||||||
"interval": 1, # Every 1 day
|
"interval": 1, # Every 1 day
|
||||||
"backupCount": 30, # Keep 30 days of logs
|
"backupCount": 30, # Keep 30 days of logs
|
||||||
"atTime": None, # Rotate at midnight (00:00)
|
"atTime": None, # Rotate at midnight (00:00)
|
||||||
"filters": ["version_filter", "request_id_filter", "security_filter"],
|
"filters": [
|
||||||
|
"version_filter",
|
||||||
|
"request_id_filter",
|
||||||
|
"client_app_filter",
|
||||||
|
"security_filter",
|
||||||
|
],
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
"loggers": {
|
"loggers": {
|
||||||
|
|||||||
+24
-10
@@ -4,6 +4,7 @@ from pathlib import Path
|
|||||||
from typing import AsyncGenerator
|
from typing import AsyncGenerator
|
||||||
|
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
|
from fastapi.exceptions import RequestValidationError
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from fastapi.responses import FileResponse, RedirectResponse
|
from fastapi.responses import FileResponse, RedirectResponse
|
||||||
from fastapi.staticfiles import StaticFiles
|
from fastapi.staticfiles import StaticFiles
|
||||||
@@ -13,10 +14,10 @@ from starlette.types import Scope
|
|||||||
|
|
||||||
from ..auth import (
|
from ..auth import (
|
||||||
periodic_dead_key_prune,
|
periodic_dead_key_prune,
|
||||||
periodic_key_reset,
|
|
||||||
periodic_stale_reservation_sweep,
|
periodic_stale_reservation_sweep,
|
||||||
)
|
)
|
||||||
from ..balance import balance_router, deprecated_wallet_router
|
from ..balance import balance_router, deprecated_wallet_router
|
||||||
|
from ..cashu_compat import install_cashu_httpx_shim
|
||||||
from ..lightning import (
|
from ..lightning import (
|
||||||
lightning_router,
|
lightning_router,
|
||||||
periodic_invoice_watcher,
|
periodic_invoice_watcher,
|
||||||
@@ -31,13 +32,18 @@ from ..nostr.discovery import providers_router
|
|||||||
from ..payment.models import models_router, update_sats_pricing
|
from ..payment.models import models_router, update_sats_pricing
|
||||||
from ..payment.price import update_prices_periodically
|
from ..payment.price import update_prices_periodically
|
||||||
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
|
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
|
||||||
|
from ..refund import periodic_refund_reconcile
|
||||||
from ..upstream.auto_topup import periodic_auto_topup
|
from ..upstream.auto_topup import periodic_auto_topup
|
||||||
from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing
|
from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing
|
||||||
from ..upstream.litellm_routing import configure_litellm
|
from ..upstream.litellm_routing import configure_litellm
|
||||||
from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout
|
from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout
|
||||||
from .admin import admin_router
|
from .admin import admin_router
|
||||||
from .db import create_session, init_db, run_migrations
|
from .db import create_session, init_db, run_migrations
|
||||||
from .exceptions import general_exception_handler, http_exception_handler
|
from .exceptions import (
|
||||||
|
general_exception_handler,
|
||||||
|
http_exception_handler,
|
||||||
|
validation_exception_handler,
|
||||||
|
)
|
||||||
from .logging import get_logger, setup_logging
|
from .logging import get_logger, setup_logging
|
||||||
from .middleware import LoggingMiddleware
|
from .middleware import LoggingMiddleware
|
||||||
from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401
|
from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401
|
||||||
@@ -63,15 +69,21 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
|||||||
models_refresh_task = None
|
models_refresh_task = None
|
||||||
model_maps_refresh_task = None
|
model_maps_refresh_task = None
|
||||||
model_paths_refresh_task = None
|
model_paths_refresh_task = None
|
||||||
key_reset_task = None
|
|
||||||
stale_reservation_task = None
|
stale_reservation_task = None
|
||||||
dead_key_prune_task = None
|
dead_key_prune_task = None
|
||||||
auto_topup_task = None
|
auto_topup_task = None
|
||||||
refund_sweep_task = None
|
refund_sweep_task = None
|
||||||
|
refund_reconcile_task = None
|
||||||
routstr_fee_task = None
|
routstr_fee_task = None
|
||||||
invoice_watcher_task = None
|
invoice_watcher_task = None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
# cashu 0.20.x passes the `proxies` kwarg httpx removed in 0.28.
|
||||||
|
# routstr.wallet and routstr.payment.lnurl also install this at import;
|
||||||
|
# repeating it here keeps startup correct for any future module that
|
||||||
|
# reaches cashu's mint client without going through those two.
|
||||||
|
install_cashu_httpx_shim()
|
||||||
|
|
||||||
# Apply litellm-wide settings (drop_params, chat-completions URL,
|
# Apply litellm-wide settings (drop_params, chat-completions URL,
|
||||||
# debug logging) before any upstream provider dispatches a request.
|
# debug logging) before any upstream provider dispatches a request.
|
||||||
configure_litellm()
|
configure_litellm()
|
||||||
@@ -143,16 +155,18 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
|||||||
refresh_model_paths_periodically(get_upstreams)
|
refresh_model_paths_periodically(get_upstreams)
|
||||||
)
|
)
|
||||||
payout_task = asyncio.create_task(periodic_payout())
|
payout_task = asyncio.create_task(periodic_payout())
|
||||||
if global_settings.nsec:
|
# Always started: the loop idles until an NSEC is configured and re-reads
|
||||||
|
# it every iteration, so a key saved (or cleared) through the admin UI
|
||||||
|
# takes effect without a restart.
|
||||||
nip91_task = asyncio.create_task(announce_provider())
|
nip91_task = asyncio.create_task(announce_provider())
|
||||||
analytics_task = asyncio.create_task(publish_usage_analytics())
|
analytics_task = asyncio.create_task(publish_usage_analytics())
|
||||||
if global_settings.providers_refresh_interval_seconds > 0:
|
if global_settings.providers_refresh_interval_seconds > 0:
|
||||||
providers_task = asyncio.create_task(providers_cache_refresher())
|
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())
|
dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune())
|
||||||
auto_topup_task = asyncio.create_task(periodic_auto_topup())
|
auto_topup_task = asyncio.create_task(periodic_auto_topup())
|
||||||
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
|
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
|
||||||
|
refund_reconcile_task = asyncio.create_task(periodic_refund_reconcile())
|
||||||
routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout())
|
routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout())
|
||||||
invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher())
|
invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher())
|
||||||
|
|
||||||
@@ -188,8 +202,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
|||||||
model_maps_refresh_task.cancel()
|
model_maps_refresh_task.cancel()
|
||||||
if model_paths_refresh_task is not None:
|
if model_paths_refresh_task is not None:
|
||||||
model_paths_refresh_task.cancel()
|
model_paths_refresh_task.cancel()
|
||||||
if key_reset_task is not None:
|
|
||||||
key_reset_task.cancel()
|
|
||||||
if stale_reservation_task is not None:
|
if stale_reservation_task is not None:
|
||||||
stale_reservation_task.cancel()
|
stale_reservation_task.cancel()
|
||||||
if dead_key_prune_task is not None:
|
if dead_key_prune_task is not None:
|
||||||
@@ -198,6 +210,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
|||||||
auto_topup_task.cancel()
|
auto_topup_task.cancel()
|
||||||
if refund_sweep_task is not None:
|
if refund_sweep_task is not None:
|
||||||
refund_sweep_task.cancel()
|
refund_sweep_task.cancel()
|
||||||
|
if refund_reconcile_task is not None:
|
||||||
|
refund_reconcile_task.cancel()
|
||||||
if routstr_fee_task is not None:
|
if routstr_fee_task is not None:
|
||||||
routstr_fee_task.cancel()
|
routstr_fee_task.cancel()
|
||||||
if invoice_watcher_task is not None:
|
if invoice_watcher_task is not None:
|
||||||
@@ -223,8 +237,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
|||||||
tasks_to_wait.append(model_maps_refresh_task)
|
tasks_to_wait.append(model_maps_refresh_task)
|
||||||
if model_paths_refresh_task is not None:
|
if model_paths_refresh_task is not None:
|
||||||
tasks_to_wait.append(model_paths_refresh_task)
|
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:
|
if stale_reservation_task is not None:
|
||||||
tasks_to_wait.append(stale_reservation_task)
|
tasks_to_wait.append(stale_reservation_task)
|
||||||
if dead_key_prune_task is not None:
|
if dead_key_prune_task is not None:
|
||||||
@@ -233,6 +245,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
|||||||
tasks_to_wait.append(auto_topup_task)
|
tasks_to_wait.append(auto_topup_task)
|
||||||
if refund_sweep_task is not None:
|
if refund_sweep_task is not None:
|
||||||
tasks_to_wait.append(refund_sweep_task)
|
tasks_to_wait.append(refund_sweep_task)
|
||||||
|
if refund_reconcile_task is not None:
|
||||||
|
tasks_to_wait.append(refund_reconcile_task)
|
||||||
if routstr_fee_task is not None:
|
if routstr_fee_task is not None:
|
||||||
tasks_to_wait.append(routstr_fee_task)
|
tasks_to_wait.append(routstr_fee_task)
|
||||||
if invoice_watcher_task is not None:
|
if invoice_watcher_task is not None:
|
||||||
@@ -293,6 +307,7 @@ app.add_middleware(LoggingMiddleware)
|
|||||||
|
|
||||||
# Add exception handlers
|
# Add exception handlers
|
||||||
app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore
|
app.add_exception_handler(HTTPException, http_exception_handler) # type: ignore
|
||||||
|
app.add_exception_handler(RequestValidationError, validation_exception_handler)
|
||||||
app.add_exception_handler(Exception, general_exception_handler)
|
app.add_exception_handler(Exception, general_exception_handler)
|
||||||
|
|
||||||
|
|
||||||
@@ -306,7 +321,6 @@ async def info() -> dict:
|
|||||||
"mints": global_settings.cashu_mints,
|
"mints": global_settings.cashu_mints,
|
||||||
"http_url": global_settings.http_url,
|
"http_url": global_settings.http_url,
|
||||||
"onion_url": global_settings.onion_url,
|
"onion_url": global_settings.onion_url,
|
||||||
"child_key_cost_msats": global_settings.child_key_cost,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2,8 +2,10 @@ import time
|
|||||||
import uuid
|
import uuid
|
||||||
from contextvars import ContextVar
|
from contextvars import ContextVar
|
||||||
from typing import Callable
|
from typing import Callable
|
||||||
|
from urllib.parse import urlsplit
|
||||||
|
|
||||||
from fastapi import Request, Response
|
from fastapi import Request, Response
|
||||||
|
from starlette.datastructures import Headers
|
||||||
from starlette.middleware.base import BaseHTTPMiddleware
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
from .logging import get_logger
|
from .logging import get_logger
|
||||||
@@ -13,6 +15,41 @@ logger = get_logger(__name__)
|
|||||||
# Context variable to store request ID across async context
|
# Context variable to store request ID across async context
|
||||||
request_id_context: ContextVar[str | None] = ContextVar("request_id")
|
request_id_context: ContextVar[str | None] = ContextVar("request_id")
|
||||||
|
|
||||||
|
client_app_context: ContextVar[str | None] = ContextVar("client_app")
|
||||||
|
|
||||||
|
UNKNOWN_CLIENT_APP = "unknown"
|
||||||
|
|
||||||
|
# Prefer OpenRouter app headers, then browser and SDK fallbacks.
|
||||||
|
_CLIENT_APP_HEADERS: tuple[str, ...] = (
|
||||||
|
"x-title",
|
||||||
|
"http-referer",
|
||||||
|
"referer",
|
||||||
|
"user-agent",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Limit untrusted header data repeated in every log record.
|
||||||
|
_CLIENT_APP_MAX_LENGTH = 120
|
||||||
|
|
||||||
|
|
||||||
|
def client_app_from_headers(headers: Headers) -> str:
|
||||||
|
for header in _CLIENT_APP_HEADERS:
|
||||||
|
raw = headers.get(header)
|
||||||
|
if raw is None:
|
||||||
|
continue
|
||||||
|
cleaned = "".join(ch for ch in raw if ch.isprintable()).strip()
|
||||||
|
if header in ("http-referer", "referer"):
|
||||||
|
try:
|
||||||
|
url = urlsplit(cleaned)
|
||||||
|
if url.scheme not in ("http", "https") or not url.hostname:
|
||||||
|
continue
|
||||||
|
except ValueError:
|
||||||
|
continue
|
||||||
|
# Attribution needs the origin, not credentials or private page URLs.
|
||||||
|
cleaned = f"{url.scheme}://{url.netloc.rsplit('@', 1)[-1]}"
|
||||||
|
if cleaned:
|
||||||
|
return cleaned[:_CLIENT_APP_MAX_LENGTH]
|
||||||
|
return UNKNOWN_CLIENT_APP
|
||||||
|
|
||||||
|
|
||||||
# Methods that are never logged: HEAD requests are health probes from
|
# Methods that are never logged: HEAD requests are health probes from
|
||||||
# monitoring/load balancers, OPTIONS are CORS preflights — both are framework
|
# monitoring/load balancers, OPTIONS are CORS preflights — both are framework
|
||||||
@@ -71,6 +108,10 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
|||||||
# Set request ID in context for logging
|
# Set request ID in context for logging
|
||||||
token = request_id_context.set(request_id)
|
token = request_id_context.set(request_id)
|
||||||
|
|
||||||
|
client_app_token = client_app_context.set(
|
||||||
|
client_app_from_headers(request.headers)
|
||||||
|
)
|
||||||
|
|
||||||
path = request.url.path
|
path = request.url.path
|
||||||
should_log = _should_log(request.method, path)
|
should_log = _should_log(request.method, path)
|
||||||
|
|
||||||
@@ -84,7 +125,9 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
|||||||
"request_id": request_id,
|
"request_id": request_id,
|
||||||
"method": request.method,
|
"method": request.method,
|
||||||
"path": path,
|
"path": path,
|
||||||
"query_params": dict(request.query_params),
|
# Names only: query values carry API keys and refund
|
||||||
|
# tokens on the wallet routes.
|
||||||
|
"query_param_names": sorted(request.query_params.keys()),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -128,6 +171,12 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
|||||||
finally:
|
finally:
|
||||||
# Reset context
|
# Reset context
|
||||||
request_id_context.reset(token)
|
request_id_context.reset(token)
|
||||||
|
client_app_context.reset(client_app_token)
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["LoggingMiddleware", "request_id_context"]
|
__all__ = [
|
||||||
|
"LoggingMiddleware",
|
||||||
|
"UNKNOWN_CLIENT_APP",
|
||||||
|
"client_app_context",
|
||||||
|
"request_id_context",
|
||||||
|
]
|
||||||
|
|||||||
+99
-16
@@ -1,8 +1,9 @@
|
|||||||
"""Redaction helpers for sensitive provider identifiers.
|
"""Redaction helpers for sensitive provider identifiers and credentials.
|
||||||
|
|
||||||
Single source of truth for stripping account-scoped identifiers (e.g. OpenAI
|
Single source of truth for stripping account-scoped identifiers (e.g. OpenAI
|
||||||
organization IDs) from any text before it is logged, returned to a caller, or
|
organization IDs) and spendable credentials (Cashu tokens, bearer keys, key
|
||||||
written to an audit entry.
|
hashes) from any text before it is logged, returned to a caller, or written to
|
||||||
|
an audit entry.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -33,19 +34,101 @@ def redact_org_ids(text: str) -> str:
|
|||||||
return _ORG_ID_PATTERN.sub(ORG_ID_PLACEHOLDER, text)
|
return _ORG_ID_PATTERN.sub(ORG_ID_PLACEHOLDER, text)
|
||||||
|
|
||||||
|
|
||||||
def redact_obj(obj: Any) -> Any:
|
SECRET_PLACEHOLDER = "[REDACTED]"
|
||||||
"""Recursively redact organization IDs in arbitrary nested structures.
|
|
||||||
|
|
||||||
Strings are redacted in place; dicts and lists/tuples are walked so that
|
# Field names whose value is spendable or authenticating on its own. Matched as
|
||||||
identifiers nested inside structured payloads (e.g. log ``extra`` fields or
|
# substrings of the lowercased key, so ``hashed_key`` (a live ``sk-`` credential)
|
||||||
error ``details``) are also stripped. Other types are returned unchanged.
|
# is stripped while the truncated ``key_hash`` prefix used for correlation is
|
||||||
"""
|
# not. Numeric values are never stripped, which keeps ``input_tokens`` and the
|
||||||
|
# other usage-analytics fields intact.
|
||||||
|
_SECRET_KEY_HINTS = (
|
||||||
|
"authorization",
|
||||||
|
"api_key",
|
||||||
|
"apikey",
|
||||||
|
"bearer",
|
||||||
|
"cashu",
|
||||||
|
"cookie",
|
||||||
|
"credential",
|
||||||
|
"hashed_key",
|
||||||
|
"mnemonic",
|
||||||
|
"nsec",
|
||||||
|
"passphrase",
|
||||||
|
"password",
|
||||||
|
"private_key",
|
||||||
|
"privkey",
|
||||||
|
"secret",
|
||||||
|
"token",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Value shapes that are spendable wherever they appear, including inside URLs,
|
||||||
|
# query strings and free-form error text. Every alternative is anchored on a
|
||||||
|
# literal prefix and uses a single bounded character class, so matching stays
|
||||||
|
# linear on the logging hot path.
|
||||||
|
_SECRET_VALUE_PATTERNS: tuple[re.Pattern[str], ...] = (
|
||||||
|
re.compile(r"cashu[A-Z][A-Za-z0-9_\-=/+]{20,}"),
|
||||||
|
re.compile(r"\bnsec1[a-z0-9]{20,}"),
|
||||||
|
re.compile(r"\bBearer\s+[A-Za-z0-9_\-.=]{10,}", re.IGNORECASE),
|
||||||
|
re.compile(r"\bsk-[A-Za-z0-9]{16,}"),
|
||||||
|
re.compile(r"\b[0-9a-f]{64}\b"),
|
||||||
|
re.compile(
|
||||||
|
r"(?<=[?&])([^=&\s]*(?:token|key|secret|password|auth|sig)[^=&\s]*=)[^&\s\"']+",
|
||||||
|
re.IGNORECASE,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
_MAX_REDACTION_DEPTH = 12
|
||||||
|
|
||||||
|
|
||||||
|
def _redact_secret_text(text: str) -> str:
|
||||||
|
for pattern in _SECRET_VALUE_PATTERNS:
|
||||||
|
text = pattern.sub(
|
||||||
|
lambda match: (match.group(1) if match.groups() else "")
|
||||||
|
+ SECRET_PLACEHOLDER,
|
||||||
|
text,
|
||||||
|
)
|
||||||
|
return redact_org_ids(text)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_secret_key(key: object) -> bool:
|
||||||
|
if not isinstance(key, str):
|
||||||
|
return False
|
||||||
|
lowered = key.lower()
|
||||||
|
return any(hint in lowered for hint in _SECRET_KEY_HINTS)
|
||||||
|
|
||||||
|
|
||||||
|
def _redact_secrets(obj: Any, depth: int, seen: frozenset[int]) -> Any:
|
||||||
if isinstance(obj, str):
|
if isinstance(obj, str):
|
||||||
return redact_org_ids(obj)
|
return _redact_secret_text(obj)
|
||||||
if isinstance(obj, dict):
|
if not isinstance(obj, (dict, list, tuple)):
|
||||||
return {key: redact_obj(value) for key, value in obj.items()}
|
|
||||||
if isinstance(obj, list):
|
|
||||||
return [redact_obj(value) for value in obj]
|
|
||||||
if isinstance(obj, tuple):
|
|
||||||
return tuple(redact_obj(value) for value in obj)
|
|
||||||
return obj
|
return obj
|
||||||
|
if depth >= _MAX_REDACTION_DEPTH or id(obj) in seen:
|
||||||
|
return SECRET_PLACEHOLDER
|
||||||
|
nested = seen | {id(obj)}
|
||||||
|
if isinstance(obj, dict):
|
||||||
|
return {
|
||||||
|
key: _redact_value(key, value, depth + 1, nested)
|
||||||
|
for key, value in obj.items()
|
||||||
|
}
|
||||||
|
redacted = [_redact_secrets(value, depth + 1, nested) for value in obj]
|
||||||
|
return redacted if isinstance(obj, list) else tuple(redacted)
|
||||||
|
|
||||||
|
|
||||||
|
def _redact_value(key: object, value: Any, depth: int, seen: frozenset[int]) -> Any:
|
||||||
|
# Containers keep being walked even under a secret-shaped key so that the
|
||||||
|
# surrounding structure stays readable for operators.
|
||||||
|
if isinstance(value, (bool, int, float, dict, list, tuple)) or value is None:
|
||||||
|
return _redact_secrets(value, depth, seen)
|
||||||
|
if _is_secret_key(key):
|
||||||
|
return SECRET_PLACEHOLDER
|
||||||
|
return _redact_secrets(value, depth, seen)
|
||||||
|
|
||||||
|
|
||||||
|
def redact_field(key: str, value: Any) -> Any:
|
||||||
|
"""Strip credentials from one named field, e.g. a log ``extra`` entry.
|
||||||
|
|
||||||
|
Both secret-shaped keys and secret-shaped values are stripped, so a leak
|
||||||
|
survives neither a renamed field nor a credential embedded in free text.
|
||||||
|
The walk is depth-limited and cycle-safe: a malformed payload degrades to
|
||||||
|
``[REDACTED]`` rather than taking the logging call down with it.
|
||||||
|
"""
|
||||||
|
return _redact_value(key, value, 0, frozenset())
|
||||||
|
|||||||
@@ -77,7 +77,6 @@ class Settings(BaseSettings):
|
|||||||
exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE")
|
exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE")
|
||||||
upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE")
|
upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE")
|
||||||
tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE")
|
tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE")
|
||||||
child_key_cost: int = Field(default=0, env="CHILD_KEY_COST")
|
|
||||||
# Minimum per-request charge in millisatoshis when model pricing is free/zero
|
# Minimum per-request charge in millisatoshis when model pricing is free/zero
|
||||||
min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT")
|
min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT")
|
||||||
reset_reserved_balance_on_startup: bool = Field(
|
reset_reserved_balance_on_startup: bool = Field(
|
||||||
@@ -101,6 +100,12 @@ class Settings(BaseSettings):
|
|||||||
|
|
||||||
# Network
|
# Network
|
||||||
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
|
cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS")
|
||||||
|
# Comma-separated METHOD:path pairs adding to the proxy's canonical
|
||||||
|
# endpoint allowlist, e.g. "POST:v1/rerank,GET:batches". Only for upstreams
|
||||||
|
# exposing an endpoint outside the OpenAI-compatible set; each addition
|
||||||
|
# widens what the provider credential can be spent against, so wildcards
|
||||||
|
# and prefixes are not supported.
|
||||||
|
proxy_extra_allowed_paths: str = Field(default="", env="PROXY_EXTRA_ALLOWED_PATHS")
|
||||||
tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL")
|
tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL")
|
||||||
providers_refresh_interval_seconds: int = Field(
|
providers_refresh_interval_seconds: int = Field(
|
||||||
default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
|
default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
|
||||||
@@ -119,7 +124,6 @@ class Settings(BaseSettings):
|
|||||||
enable_model_paths_refresh: bool = Field(
|
enable_model_paths_refresh: bool = Field(
|
||||||
default=True, env="ENABLE_MODEL_PATHS_REFRESH"
|
default=True, env="ENABLE_MODEL_PATHS_REFRESH"
|
||||||
)
|
)
|
||||||
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
|
|
||||||
# Uncollected refund tokens are swept after ~6 months (180 days).
|
# Uncollected refund tokens are swept after ~6 months (180 days).
|
||||||
# Fixed for now: not configurable via env or the settings DB/admin API
|
# Fixed for now: not configurable via env or the settings DB/admin API
|
||||||
# (empty env list disables env binding; see FIXED_FIELDS).
|
# (empty env list disables env binding; see FIXED_FIELDS).
|
||||||
@@ -127,6 +131,14 @@ class Settings(BaseSettings):
|
|||||||
refund_sweep_claim_timeout_seconds: int = Field(
|
refund_sweep_claim_timeout_seconds: int = Field(
|
||||||
default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS"
|
default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS"
|
||||||
)
|
)
|
||||||
|
# How long an open refund claim may sit before the reconciler asks the mint
|
||||||
|
# what became of it. Doubles as the reconciler's per-row lease.
|
||||||
|
refund_claim_timeout_seconds: int = Field(
|
||||||
|
default=300, gt=0, env="REFUND_CLAIM_TIMEOUT_SECONDS"
|
||||||
|
)
|
||||||
|
refund_reconcile_interval_seconds: int = Field(
|
||||||
|
default=60, gt=0, env="REFUND_RECONCILE_INTERVAL_SECONDS"
|
||||||
|
)
|
||||||
|
|
||||||
# Database connection-pool controls (advanced). Capacity defaults provide
|
# Database connection-pool controls (advanced). Capacity defaults provide
|
||||||
# headroom for Routstr's concurrent request and background-payment workload.
|
# headroom for Routstr's concurrent request and background-payment workload.
|
||||||
@@ -142,6 +154,9 @@ class Settings(BaseSettings):
|
|||||||
database_pool_hold_warn_seconds: float = Field(
|
database_pool_hold_warn_seconds: float = Field(
|
||||||
default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS"
|
default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS"
|
||||||
)
|
)
|
||||||
|
database_busy_timeout: float = Field(
|
||||||
|
default=30.0, gt=0, env="DATABASE_BUSY_TIMEOUT"
|
||||||
|
)
|
||||||
|
|
||||||
# Logging
|
# Logging
|
||||||
log_level: str = Field(default="INFO", env="LOG_LEVEL")
|
log_level: str = Field(default="INFO", env="LOG_LEVEL")
|
||||||
@@ -192,12 +207,18 @@ SECRET_FIELDS = frozenset({"admin_password", "nsec"})
|
|||||||
# neither store nor shadow them; env is always authoritative.
|
# neither store nor shadow them; env is always authoritative.
|
||||||
ENV_ONLY_FIELDS = frozenset(
|
ENV_ONLY_FIELDS = frozenset(
|
||||||
{
|
{
|
||||||
|
# Widening the proxy's reachable upstream surface is a deployment
|
||||||
|
# decision, not a runtime toggle: it changes what the provider
|
||||||
|
# credential can be spent against. Keeping it env-only also lets the
|
||||||
|
# proxy parse it once at import without going stale.
|
||||||
|
"proxy_extra_allowed_paths",
|
||||||
"database_pool_size",
|
"database_pool_size",
|
||||||
"database_max_overflow",
|
"database_max_overflow",
|
||||||
"database_pool_timeout",
|
"database_pool_timeout",
|
||||||
"database_pool_recycle",
|
"database_pool_recycle",
|
||||||
"database_pool_pre_ping",
|
"database_pool_pre_ping",
|
||||||
"database_pool_hold_warn_seconds",
|
"database_pool_hold_warn_seconds",
|
||||||
|
"database_busy_timeout",
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -248,7 +269,7 @@ def derive_npub_from_nsec(nsec: str) -> str | None:
|
|||||||
boot.
|
boot.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
from nostr.key import PublicKey # type: ignore
|
from nostr_sdk import PublicKey
|
||||||
|
|
||||||
from ..nostr.listing import nsec_to_keypair
|
from ..nostr.listing import nsec_to_keypair
|
||||||
except ImportError:
|
except ImportError:
|
||||||
@@ -260,7 +281,7 @@ def derive_npub_from_nsec(nsec: str) -> str | None:
|
|||||||
_privkey_hex, pubkey_hex = keypair
|
_privkey_hex, pubkey_hex = keypair
|
||||||
|
|
||||||
try:
|
try:
|
||||||
return PublicKey(bytes.fromhex(pubkey_hex)).bech32()
|
return PublicKey.parse(pubkey_hex).to_bech32()
|
||||||
except (ValueError, AttributeError):
|
except (ValueError, AttributeError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ import subprocess
|
|||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
BASE_VERSION = "0.4.5"
|
BASE_VERSION = "0.4.7"
|
||||||
|
|
||||||
_REPO_ROOT = Path(__file__).resolve().parents[2]
|
_REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||||
_GIT_TIMEOUT_SECONDS = 2.0
|
_GIT_TIMEOUT_SECONDS = 2.0
|
||||||
|
|||||||
+9
-12
@@ -80,8 +80,6 @@ class _InvoiceSettlement:
|
|||||||
purpose: str
|
purpose: str
|
||||||
api_key_hash: str | None
|
api_key_hash: str | None
|
||||||
mint_url: str | None
|
mint_url: str | None
|
||||||
balance_limit: int | None
|
|
||||||
balance_limit_reset: str | None
|
|
||||||
validity_date: int | None
|
validity_date: int | None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -93,8 +91,6 @@ class _InvoiceSettlement:
|
|||||||
purpose=invoice.purpose,
|
purpose=invoice.purpose,
|
||||||
api_key_hash=invoice.api_key_hash,
|
api_key_hash=invoice.api_key_hash,
|
||||||
mint_url=invoice.mint_url,
|
mint_url=invoice.mint_url,
|
||||||
balance_limit=invoice.balance_limit,
|
|
||||||
balance_limit_reset=invoice.balance_limit_reset,
|
|
||||||
validity_date=invoice.validity_date,
|
validity_date=invoice.validity_date,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -118,8 +114,6 @@ class InvoiceCreateRequest(BaseModel):
|
|||||||
default=None,
|
default=None,
|
||||||
description="Deprecated: legacy field for topup. Prefer Authorization header.",
|
description="Deprecated: legacy field for topup. Prefer Authorization header.",
|
||||||
)
|
)
|
||||||
balance_limit: int | None = Field(default=None)
|
|
||||||
balance_limit_reset: str | None = Field(default=None)
|
|
||||||
validity_date: int | None = Field(default=None)
|
validity_date: int | None = Field(default=None)
|
||||||
|
|
||||||
|
|
||||||
@@ -220,11 +214,18 @@ async def _request_mint_with_fallback(
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False)
|
wallet = await get_wallet(
|
||||||
|
mint_url,
|
||||||
|
"sat",
|
||||||
|
retry_on_rate_limit=False,
|
||||||
|
load_proofs=False,
|
||||||
|
)
|
||||||
quote = await run_mint_operation(
|
quote = await run_mint_operation(
|
||||||
lambda: wallet.request_mint(amount_sats),
|
lambda: wallet.request_mint(amount_sats),
|
||||||
op_name="request_mint_invoice",
|
op_name="request_mint_invoice",
|
||||||
mint_url=mint_url,
|
mint_url=mint_url,
|
||||||
|
# Quote creation is unsafe to retry without idempotency.
|
||||||
|
retry_timeouts=False,
|
||||||
retry_on_rate_limit=False,
|
retry_on_rate_limit=False,
|
||||||
)
|
)
|
||||||
return quote.request, quote.quote, mint_url
|
return quote.request, quote.quote, mint_url
|
||||||
@@ -400,8 +401,6 @@ async def create_invoice(
|
|||||||
api_key_hash=api_key_token[3:] if api_key_token else None,
|
api_key_hash=api_key_token[3:] if api_key_token else None,
|
||||||
purpose=request.purpose,
|
purpose=request.purpose,
|
||||||
mint_url=mint_url,
|
mint_url=mint_url,
|
||||||
balance_limit=request.balance_limit,
|
|
||||||
balance_limit_reset=request.balance_limit_reset,
|
|
||||||
validity_date=request.validity_date,
|
validity_date=request.validity_date,
|
||||||
expires_at=expires_at,
|
expires_at=expires_at,
|
||||||
)
|
)
|
||||||
@@ -582,7 +581,7 @@ async def check_invoice_payment(
|
|||||||
await session.commit()
|
await session.commit()
|
||||||
|
|
||||||
mint_url = settlement.mint_url or settings.primary_mint
|
mint_url = settlement.mint_url or settings.primary_mint
|
||||||
wallet = await get_wallet(mint_url, "sat")
|
wallet = await get_wallet(mint_url, "sat", load_proofs=False)
|
||||||
try:
|
try:
|
||||||
mint_status = await run_mint_operation(
|
mint_status = await run_mint_operation(
|
||||||
lambda: wallet.get_mint_quote(settlement.payment_hash),
|
lambda: wallet.get_mint_quote(settlement.payment_hash),
|
||||||
@@ -797,8 +796,6 @@ async def _create_api_key_record(
|
|||||||
balance=invoice.amount_sats * 1000,
|
balance=invoice.amount_sats * 1000,
|
||||||
refund_currency="sat",
|
refund_currency="sat",
|
||||||
refund_mint_url=mint_url,
|
refund_mint_url=mint_url,
|
||||||
balance_limit=invoice.balance_limit,
|
|
||||||
balance_limit_reset=invoice.balance_limit_reset,
|
|
||||||
validity_date=invoice.validity_date,
|
validity_date=invoice.validity_date,
|
||||||
)
|
)
|
||||||
session.add(api_key)
|
session.add(api_key)
|
||||||
|
|||||||
+24
-1
@@ -177,6 +177,8 @@ class MintRateGuard:
|
|||||||
if isinstance(error, httpx.HTTPStatusError):
|
if isinstance(error, httpx.HTTPStatusError):
|
||||||
retry_after = parse_retry_after(error.response.headers)
|
retry_after = parse_retry_after(error.response.headers)
|
||||||
self.apply_rate_limit_cooldown(retry_after)
|
self.apply_rate_limit_cooldown(retry_after)
|
||||||
|
elif is_mint_transport_error(error):
|
||||||
|
self.apply_cooldown(MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport")
|
||||||
else:
|
else:
|
||||||
self.apply_cooldown(1.0)
|
self.apply_cooldown(1.0)
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -236,6 +238,17 @@ def mint_cooldown_reason(mint_url: str) -> str | None:
|
|||||||
return MintRateGuard.get(mint_url).cooldown_reason()
|
return MintRateGuard.get(mint_url).cooldown_reason()
|
||||||
|
|
||||||
|
|
||||||
|
def is_mint_transport_error(error: BaseException) -> bool:
|
||||||
|
current: BaseException | None = error
|
||||||
|
seen: set[int] = set()
|
||||||
|
while current is not None and id(current) not in seen:
|
||||||
|
seen.add(id(current))
|
||||||
|
if isinstance(current, MINT_TRANSPORT_EXCEPTIONS):
|
||||||
|
return True
|
||||||
|
current = current.__cause__ or current.__context__
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def is_mint_rate_limited(error: BaseException) -> bool:
|
def is_mint_rate_limited(error: BaseException) -> bool:
|
||||||
"""Return whether an exception chain represents HTTP 429/cooldown."""
|
"""Return whether an exception chain represents HTTP 429/cooldown."""
|
||||||
|
|
||||||
@@ -269,6 +282,7 @@ async def run_mint_operation(
|
|||||||
mint_url: str = "",
|
mint_url: str = "",
|
||||||
retry_timeouts: bool = True,
|
retry_timeouts: bool = True,
|
||||||
retry_on_rate_limit: bool = True,
|
retry_on_rate_limit: bool = True,
|
||||||
|
allow_during_cooldown: bool = False,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""Run one mint operation with bounded concurrency and adaptive cooldown."""
|
"""Run one mint operation with bounded concurrency and adaptive cooldown."""
|
||||||
|
|
||||||
@@ -282,7 +296,7 @@ async def run_mint_operation(
|
|||||||
return await factory()
|
return await factory()
|
||||||
|
|
||||||
async def invoke() -> Any:
|
async def invoke() -> Any:
|
||||||
if guard is not None:
|
if guard is not None and not allow_during_cooldown:
|
||||||
return await guard.run(timed_factory)
|
return await guard.run(timed_factory)
|
||||||
return await timed_factory()
|
return await timed_factory()
|
||||||
|
|
||||||
@@ -293,6 +307,7 @@ async def run_mint_operation(
|
|||||||
raise
|
raise
|
||||||
except (asyncio.TimeoutError, httpx.TimeoutException) as exc:
|
except (asyncio.TimeoutError, httpx.TimeoutException) as exc:
|
||||||
if retry_timeouts and attempt < max_attempts - 1:
|
if retry_timeouts and attempt < max_attempts - 1:
|
||||||
|
# Cooldown opens only after retries; earlier would stretch each backoff to a full cooldown wait.
|
||||||
backoff = (2**attempt) + (time.monotonic() % 1.0)
|
backoff = (2**attempt) + (time.monotonic() % 1.0)
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Mint operation timed out, retrying",
|
"Mint operation timed out, retrying",
|
||||||
@@ -305,11 +320,19 @@ async def run_mint_operation(
|
|||||||
)
|
)
|
||||||
await asyncio.sleep(backoff)
|
await asyncio.sleep(backoff)
|
||||||
continue
|
continue
|
||||||
|
if guard is not None:
|
||||||
|
guard.apply_cooldown(
|
||||||
|
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
|
||||||
|
)
|
||||||
raise httpx.TimeoutException(
|
raise httpx.TimeoutException(
|
||||||
f"{op_name} timed out (attempts: {attempt + 1})"
|
f"{op_name} timed out (attempts: {attempt + 1})"
|
||||||
) from exc
|
) from exc
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
if not is_mint_rate_limited(exc):
|
if not is_mint_rate_limited(exc):
|
||||||
|
if guard is not None and is_mint_transport_error(exc):
|
||||||
|
guard.apply_cooldown(
|
||||||
|
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
|
||||||
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
backoff = (2**attempt) + (time.monotonic() % 1.0)
|
backoff = (2**attempt) + (time.monotonic() % 1.0)
|
||||||
|
|||||||
@@ -12,13 +12,11 @@ import json
|
|||||||
import time
|
import time
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from nostr.event import Event
|
|
||||||
from nostr.key import PrivateKey
|
|
||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.log_manager import log_manager
|
from ..core.log_manager import log_manager
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
from .listing import nsec_to_keypair, publish_to_relay
|
from .listing import nsec_to_keypair, publish_to_relay
|
||||||
|
from .sdk import create_signed_event
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -44,18 +42,6 @@ WINDOW_DEFINITIONS: tuple[tuple[str, int, int], ...] = (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _event_to_dict(ev: Event) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"id": ev.id,
|
|
||||||
"pubkey": ev.public_key,
|
|
||||||
"created_at": ev.created_at,
|
|
||||||
"kind": int(ev.kind) if not isinstance(ev.kind, int) else ev.kind,
|
|
||||||
"tags": ev.tags,
|
|
||||||
"content": ev.content,
|
|
||||||
"sig": ev.signature,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve_provider_id(public_key_hex: str) -> str:
|
def _resolve_provider_id(public_key_hex: str) -> str:
|
||||||
explicit_provider_id = (settings.provider_id or "").strip()
|
explicit_provider_id = (settings.provider_id or "").strip()
|
||||||
if explicit_provider_id:
|
if explicit_provider_id:
|
||||||
@@ -293,21 +279,18 @@ def create_stats_snapshot_event(
|
|||||||
*,
|
*,
|
||||||
d_tag: str,
|
d_tag: str,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
private_key = PrivateKey(bytes.fromhex(private_key_hex))
|
|
||||||
tags = [
|
tags = [
|
||||||
["d", d_tag],
|
["d", d_tag],
|
||||||
["provider", provider_id],
|
["provider", provider_id],
|
||||||
["schema", ANALYTICS_SCHEMA],
|
["schema", ANALYTICS_SCHEMA],
|
||||||
]
|
]
|
||||||
|
|
||||||
event = Event(
|
return create_signed_event(
|
||||||
public_key=private_key.public_key.hex(),
|
private_key_hex,
|
||||||
content=payload_json,
|
|
||||||
kind=ANALYTICS_KIND,
|
kind=ANALYTICS_KIND,
|
||||||
|
content=payload_json,
|
||||||
tags=tags,
|
tags=tags,
|
||||||
)
|
)
|
||||||
private_key.sign_event(event)
|
|
||||||
return _event_to_dict(event)
|
|
||||||
|
|
||||||
|
|
||||||
def _fingerprint_payload(payload: dict[str, Any]) -> str:
|
def _fingerprint_payload(payload: dict[str, Any]) -> str:
|
||||||
|
|||||||
@@ -320,18 +320,18 @@ async def fetch_provider_health(endpoint_url: str) -> dict[str, Any]:
|
|||||||
is_onion = ".onion" in endpoint_url
|
is_onion = ".onion" in endpoint_url
|
||||||
|
|
||||||
# Set up client arguments conditionally
|
# Set up client arguments conditionally
|
||||||
proxies = None
|
proxy: str | None = None
|
||||||
if is_onion:
|
if is_onion:
|
||||||
try:
|
try:
|
||||||
tor_proxy = settings.tor_proxy_url
|
tor_proxy = settings.tor_proxy_url
|
||||||
except Exception:
|
except Exception:
|
||||||
tor_proxy = "socks5://127.0.0.1:9050"
|
tor_proxy = "socks5://127.0.0.1:9050"
|
||||||
proxies = {"http://": tor_proxy, "https://": tor_proxy} # type: ignore[assignment]
|
proxy = tor_proxy
|
||||||
|
|
||||||
async with httpx.AsyncClient(
|
async with httpx.AsyncClient(
|
||||||
timeout=httpx.Timeout(30.0),
|
timeout=httpx.Timeout(30.0),
|
||||||
follow_redirects=True,
|
follow_redirects=True,
|
||||||
proxies=proxies, # type: ignore[arg-type]
|
proxy=proxy,
|
||||||
) as client:
|
) as client:
|
||||||
# Prefer provider's /v1/info for full details
|
# Prefer provider's /v1/info for full details
|
||||||
info_url = f"{endpoint_url.rstrip('/')}/v1/info"
|
info_url = f"{endpoint_url.rstrip('/')}/v1/info"
|
||||||
|
|||||||
+168
-239
@@ -8,18 +8,12 @@ import asyncio
|
|||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import random
|
import random
|
||||||
import ssl
|
|
||||||
import time
|
import time
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
from nostr.event import Event
|
|
||||||
from nostr.filter import Filter, Filters
|
|
||||||
from nostr.key import PrivateKey
|
|
||||||
from nostr.message_type import ClientMessageType
|
|
||||||
from nostr.relay_manager import RelayManager
|
|
||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
|
from .sdk import create_signed_event, fetch_events, parse_keypair, send_event
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -33,18 +27,6 @@ def get_app_version() -> str | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _event_to_dict(ev: Event) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"id": ev.id,
|
|
||||||
"pubkey": ev.public_key,
|
|
||||||
"created_at": ev.created_at,
|
|
||||||
"kind": int(ev.kind) if not isinstance(ev.kind, int) else ev.kind,
|
|
||||||
"tags": ev.tags,
|
|
||||||
"content": ev.content,
|
|
||||||
"sig": ev.signature,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def nsec_to_keypair(nsec: str) -> tuple[str, str] | None:
|
def nsec_to_keypair(nsec: str) -> tuple[str, str] | None:
|
||||||
"""
|
"""
|
||||||
Convert a Nostr private key (nsec) to a keypair (privkey_hex, pubkey_hex).
|
Convert a Nostr private key (nsec) to a keypair (privkey_hex, pubkey_hex).
|
||||||
@@ -56,16 +38,10 @@ def nsec_to_keypair(nsec: str) -> tuple[str, str] | None:
|
|||||||
Tuple of (private_key_hex, public_key_hex) or None if invalid
|
Tuple of (private_key_hex, public_key_hex) or None if invalid
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
if nsec.startswith("nsec"):
|
if not (nsec.startswith("nsec") or len(nsec) == 64):
|
||||||
pk = PrivateKey.from_nsec(nsec)
|
|
||||||
return (pk.hex(), pk.public_key.hex())
|
|
||||||
|
|
||||||
if len(nsec) == 64:
|
|
||||||
pk = PrivateKey(bytes.fromhex(nsec))
|
|
||||||
return (pk.hex(), pk.public_key.hex())
|
|
||||||
|
|
||||||
logger.error(f"Invalid private key format/length: {len(nsec)}")
|
logger.error(f"Invalid private key format/length: {len(nsec)}")
|
||||||
return None
|
return None
|
||||||
|
return parse_keypair(nsec)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to convert nsec to keypair: {e}")
|
logger.error(f"Failed to convert nsec to keypair: {e}")
|
||||||
return None
|
return None
|
||||||
@@ -93,8 +69,6 @@ def create_listing_event(
|
|||||||
Returns:
|
Returns:
|
||||||
Complete signed nostr event as a dict ready for publishing
|
Complete signed nostr event as a dict ready for publishing
|
||||||
"""
|
"""
|
||||||
pk = PrivateKey(bytes.fromhex(private_key_hex))
|
|
||||||
|
|
||||||
tags = [["d", provider_id]]
|
tags = [["d", provider_id]]
|
||||||
for url in endpoint_urls:
|
for url in endpoint_urls:
|
||||||
tags.append(["u", url])
|
tags.append(["u", url])
|
||||||
@@ -107,9 +81,12 @@ def create_listing_event(
|
|||||||
|
|
||||||
content = json.dumps(metadata, separators=(",", ":")) if metadata else ""
|
content = json.dumps(metadata, separators=(",", ":")) if metadata else ""
|
||||||
|
|
||||||
ev = Event(pk.public_key.hex(), content, kind=38421, tags=tags)
|
return create_signed_event(
|
||||||
pk.sign_event(ev)
|
private_key_hex,
|
||||||
return _event_to_dict(ev)
|
kind=38421,
|
||||||
|
content=content,
|
||||||
|
tags=tags,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _get_tag_values(event: dict[str, Any], key: str) -> list[str]:
|
def _get_tag_values(event: dict[str, Any], key: str) -> list[str]:
|
||||||
@@ -177,75 +154,25 @@ async def query_listing_events(
|
|||||||
succeeded without transport-level errors.
|
succeeded without transport-level errors.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def _sync_query() -> tuple[list[dict[str, Any]], bool]:
|
|
||||||
rm = RelayManager()
|
|
||||||
rm.add_relay(relay_url)
|
|
||||||
events_out: list[dict[str, Any]] = []
|
|
||||||
ok = True
|
|
||||||
try:
|
try:
|
||||||
rm.open_connections({"cert_reqs": ssl.CERT_NONE})
|
events_out = await fetch_events(
|
||||||
time.sleep(1.0)
|
relay_url,
|
||||||
|
kind=38421,
|
||||||
flt = Filter(kinds=[38421], authors=[pubkey], limit=10)
|
author=pubkey,
|
||||||
filters = Filters([flt])
|
limit=10,
|
||||||
sub_id = f"routstr_listing_{int(time.time())}"
|
timeout=timeout,
|
||||||
rm.add_subscription(sub_id, filters)
|
|
||||||
req: list[Any] = [ClientMessageType.REQUEST, sub_id]
|
|
||||||
req.extend(filters.to_json_array())
|
|
||||||
rm.publish_message(json.dumps(req))
|
|
||||||
|
|
||||||
start = time.time()
|
|
||||||
last_event_ts = start
|
|
||||||
while time.time() - start < timeout:
|
|
||||||
drained = False
|
|
||||||
while rm.message_pool.has_events():
|
|
||||||
drained = True
|
|
||||||
ev_msg = rm.message_pool.get_event()
|
|
||||||
ev = ev_msg.event
|
|
||||||
ev_dict = _event_to_dict(ev)
|
|
||||||
if provider_id is not None:
|
|
||||||
tags = ev_dict.get("tags", [])
|
|
||||||
if not any(
|
|
||||||
isinstance(t, list)
|
|
||||||
and len(t) >= 2
|
|
||||||
and t[0] == "d"
|
|
||||||
and t[1] == provider_id
|
|
||||||
for t in tags
|
|
||||||
):
|
|
||||||
continue
|
|
||||||
events_out.append(ev_dict)
|
|
||||||
logger.debug(
|
|
||||||
f"Found listing event: {ev_dict.get('id', '')[:6]}...{ev_dict.get('id', '')[-6:]}"
|
|
||||||
)
|
)
|
||||||
if drained:
|
|
||||||
last_event_ts = time.time()
|
|
||||||
|
|
||||||
while rm.message_pool.has_notices():
|
|
||||||
notice = rm.message_pool.get_notice()
|
|
||||||
try:
|
|
||||||
content = getattr(notice, "content", notice)
|
|
||||||
s = str(content)
|
|
||||||
if len(s) > 200:
|
|
||||||
s = s[:200] + "..."
|
|
||||||
logger.debug(f"Relay notice: {s}")
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
if time.time() - last_event_ts > 2.5:
|
|
||||||
break
|
|
||||||
|
|
||||||
time.sleep(0.1)
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
ok = False
|
|
||||||
logger.debug(f"Failed to query relay {relay_url}: {type(e).__name__}")
|
logger.debug(f"Failed to query relay {relay_url}: {type(e).__name__}")
|
||||||
finally:
|
return [], False
|
||||||
try:
|
|
||||||
rm.close_connections()
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
return events_out, ok
|
|
||||||
|
|
||||||
return await asyncio.to_thread(_sync_query)
|
if provider_id is not None:
|
||||||
|
events_out = [
|
||||||
|
event
|
||||||
|
for event in events_out
|
||||||
|
if _get_single_tag_value(event, "d") == provider_id
|
||||||
|
]
|
||||||
|
return events_out, True
|
||||||
|
|
||||||
|
|
||||||
def discover_onion_url_from_tor(base_dir: str = "/var/lib/tor") -> str | None:
|
def discover_onion_url_from_tor(base_dir: str = "/var/lib/tor") -> str | None:
|
||||||
@@ -333,117 +260,102 @@ async def publish_to_relay(
|
|||||||
Publish a listing event to a nostr relay via nostr library.
|
Publish a listing event to a nostr relay via nostr library.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def _sync_publish() -> bool:
|
|
||||||
rm = RelayManager()
|
|
||||||
rm.add_relay(relay_url)
|
|
||||||
try:
|
try:
|
||||||
rm.open_connections({"cert_reqs": ssl.CERT_NONE})
|
await send_event(relay_url, event, timeout=timeout)
|
||||||
time.sleep(1.0)
|
|
||||||
# Publish the event as-is via publish_message to preserve signature
|
|
||||||
rm.publish_message(json.dumps(["EVENT", event]))
|
|
||||||
logger.debug(f"Sent listing event {event.get('id', '')} to {relay_url}")
|
logger.debug(f"Sent listing event {event.get('id', '')} to {relay_url}")
|
||||||
time.sleep(1.0)
|
|
||||||
return True
|
return True
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug(f"Failed to publish to {relay_url}: {type(e).__name__}")
|
logger.debug(f"Failed to publish to {relay_url}: {type(e).__name__}")
|
||||||
return False
|
return False
|
||||||
finally:
|
|
||||||
try:
|
|
||||||
rm.close_connections()
|
|
||||||
except Exception:
|
|
||||||
pass
|
|
||||||
|
|
||||||
return await asyncio.to_thread(_sync_publish)
|
|
||||||
|
|
||||||
|
|
||||||
async def announce_provider() -> None:
|
# Re-announce cadence once a provider is listed.
|
||||||
"""
|
ANNOUNCEMENT_INTERVAL_SECONDS = 24 * 60 * 60
|
||||||
Background task to announce this Routstr provider to Nostr relays.
|
# Poll cadence while there is nothing to announce (no NSEC, no endpoint, ...).
|
||||||
Checks for existing announcements and creates new ones if needed.
|
DISABLED_POLL_SECONDS = 60
|
||||||
"""
|
# How often the long re-announce sleep re-checks the configured NSEC, so a
|
||||||
# Check for NSEC in environment (use NSEC only)
|
# newly saved identity is announced promptly instead of up to 24h later.
|
||||||
nsec = settings.nsec
|
IDENTITY_POLL_SECONDS = 30
|
||||||
if not nsec:
|
|
||||||
logger.info("Nostr private key not found (NSEC), skipping listing announcement")
|
|
||||||
return
|
|
||||||
|
|
||||||
# Convert NSEC to keypair
|
DEFAULT_RELAY_URLS = [
|
||||||
keypair = nsec_to_keypair(nsec)
|
"wss://relay.nostr.band",
|
||||||
if not keypair:
|
"wss://relay.damus.io",
|
||||||
logger.error("Failed to parse NSEC, skipping listing announcement")
|
"wss://relay.routstr.com",
|
||||||
return
|
"wss://nos.lol",
|
||||||
|
]
|
||||||
|
|
||||||
private_key_hex, public_key_hex = keypair
|
|
||||||
logger.info(f"Using Nostr pubkey: {public_key_hex}")
|
|
||||||
|
|
||||||
# Resolve settings and determine if we can publish BEFORE touching relays
|
def _resolve_endpoint_urls() -> list[str]:
|
||||||
try:
|
"""Endpoints to advertise: a public HTTP URL and/or an onion URL."""
|
||||||
base_url: str | None = settings.http_url
|
endpoint_urls: list[str] = []
|
||||||
onion_url: str | None = settings.onion_url
|
|
||||||
provider_name = settings.name or "Routstr Proxy"
|
base_url = (settings.http_url or "").strip()
|
||||||
provider_about = settings.description or "Privacy-preserving AI proxy via Nostr"
|
if base_url and base_url != "http://localhost:8000":
|
||||||
cashu_mints = [m.strip() for m in settings.cashu_mints if m.strip()]
|
endpoint_urls.append(base_url)
|
||||||
except Exception:
|
|
||||||
base_url = settings.http_url or None
|
onion_url = (settings.onion_url or "").strip()
|
||||||
onion_url = settings.onion_url or None
|
|
||||||
provider_name = settings.name or "Routstr Proxy"
|
|
||||||
provider_about = settings.description or "Privacy-preserving AI proxy via Nostr"
|
|
||||||
cashu_mints = [m.strip() for m in settings.cashu_mints if m.strip()]
|
|
||||||
if not onion_url:
|
if not onion_url:
|
||||||
discovered = discover_onion_url_from_tor()
|
discovered = discover_onion_url_from_tor()
|
||||||
if discovered:
|
if discovered:
|
||||||
onion_url = discovered
|
onion_url = discovered
|
||||||
logger.info(f"Discovered onion URL via Tor volume: {onion_url}")
|
logger.info(f"Discovered onion URL via Tor volume: {onion_url}")
|
||||||
mint_urls = cashu_mints if cashu_mints else None
|
|
||||||
|
|
||||||
endpoint_urls: list[str] = []
|
if onion_url:
|
||||||
if base_url and base_url.strip() and base_url.strip() != "http://localhost:8000":
|
if onion_url.endswith(".onion") and not (
|
||||||
endpoint_urls.append(base_url.strip())
|
onion_url.startswith("http://") or onion_url.startswith("https://")
|
||||||
if onion_url and onion_url.strip():
|
|
||||||
ou = onion_url.strip()
|
|
||||||
if ou.endswith(".onion") and not (
|
|
||||||
ou.startswith("http://") or ou.startswith("https://")
|
|
||||||
):
|
):
|
||||||
ou = f"http://{ou}"
|
onion_url = f"http://{onion_url}"
|
||||||
endpoint_urls.append(ou)
|
endpoint_urls.append(onion_url)
|
||||||
|
|
||||||
if not endpoint_urls:
|
return endpoint_urls
|
||||||
logger.warning(
|
|
||||||
"No valid endpoints configured (HTTP_URL/ONION_URL). Skipping listing publish."
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
# Only now configure relays and determine provider_id (may query relays)
|
|
||||||
|
def _resolve_relay_urls() -> list[str]:
|
||||||
relay_urls = [u.strip() for u in getattr(settings, "relays", []) if u.strip()]
|
relay_urls = [u.strip() for u in getattr(settings, "relays", []) if u.strip()]
|
||||||
if not relay_urls:
|
return relay_urls or list(DEFAULT_RELAY_URLS)
|
||||||
relay_urls = [
|
|
||||||
"wss://relay.nostr.band",
|
|
||||||
"wss://relay.damus.io",
|
|
||||||
"wss://relay.routstr.com",
|
|
||||||
"wss://nos.lol",
|
|
||||||
]
|
|
||||||
|
|
||||||
provider_id = await _determine_provider_id(public_key_hex, relay_urls)
|
|
||||||
logger.info(f"Using provider_id: {provider_id}")
|
|
||||||
|
|
||||||
# Build metadata
|
def _resolve_mint_urls() -> list[str] | None:
|
||||||
metadata = {
|
mints = [m.strip() for m in (settings.cashu_mints or []) if m.strip()]
|
||||||
"name": provider_name,
|
return mints or None
|
||||||
"about": provider_about,
|
|
||||||
}
|
|
||||||
|
|
||||||
# Create the candidate event that we would publish
|
|
||||||
version_str = get_app_version()
|
|
||||||
candidate_event = create_listing_event(
|
|
||||||
private_key_hex=private_key_hex,
|
|
||||||
provider_id=provider_id,
|
|
||||||
endpoint_urls=endpoint_urls,
|
|
||||||
mint_urls=mint_urls,
|
|
||||||
version=version_str,
|
|
||||||
metadata=metadata,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Backoff configuration and state (sensible defaults)
|
async def _sleep_until_next_announcement(
|
||||||
|
seconds: float, parsed_nsec: str | None
|
||||||
|
) -> None:
|
||||||
|
"""Sleep up to ``seconds``, returning early if the configured NSEC changes.
|
||||||
|
|
||||||
|
Without the early wake, a node runner who replaces the NSEC in the admin UI
|
||||||
|
would wait out the whole re-announce interval before the new identity (and,
|
||||||
|
with it, the new ``d`` tag and npub) is announced.
|
||||||
|
"""
|
||||||
|
remaining = float(seconds)
|
||||||
|
while remaining > 0:
|
||||||
|
if (settings.nsec or "").strip() != (parsed_nsec or ""):
|
||||||
|
return
|
||||||
|
chunk = min(float(IDENTITY_POLL_SECONDS), remaining)
|
||||||
|
await asyncio.sleep(chunk)
|
||||||
|
remaining -= chunk
|
||||||
|
|
||||||
|
|
||||||
|
async def announce_provider() -> None:
|
||||||
|
"""Background task announcing this Routstr provider to Nostr relays.
|
||||||
|
|
||||||
|
Started unconditionally at boot: while the node has no NSEC the task idles
|
||||||
|
and re-checks, so an identity configured later through the admin UI is
|
||||||
|
picked up (and announced) without a restart. The identity, endpoints, mints
|
||||||
|
and relays are all re-read every iteration, mirroring
|
||||||
|
``publish_usage_analytics``.
|
||||||
|
"""
|
||||||
|
parsed_nsec: str | None = None
|
||||||
|
private_key_hex: str | None = None
|
||||||
|
public_key_hex: str | None = None
|
||||||
|
provider_id: str | None = None
|
||||||
|
warned_missing_nsec = False
|
||||||
|
|
||||||
|
# Backoff state is deliberately long-lived: it has to survive an idle poll,
|
||||||
|
# a full re-announce cycle and an identity change, otherwise a failing relay
|
||||||
|
# would be retried at full rate on every pass.
|
||||||
backoff_base = 5.0
|
backoff_base = 5.0
|
||||||
backoff_max = 900.0
|
backoff_max = 900.0
|
||||||
backoff_jitter_ratio = 0.2
|
backoff_jitter_ratio = 0.2
|
||||||
@@ -468,67 +380,74 @@ async def announce_provider() -> None:
|
|||||||
f"Backoff: {relay} delay={delay:.1f}s jitter={jitter:.1f}s next={int(scheduled)}"
|
f"Backoff: {relay} delay={delay:.1f}s jitter={jitter:.1f}s next={int(scheduled)}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Fetch existing events for this provider_id
|
|
||||||
existing_events: list[dict[str, Any]] = []
|
|
||||||
for relay_url in relay_urls:
|
|
||||||
if _should_skip(relay_url):
|
|
||||||
logger.debug(f"Skipping {relay_url} due to backoff")
|
|
||||||
continue
|
|
||||||
events, ok = await query_listing_events(relay_url, public_key_hex, provider_id)
|
|
||||||
if ok:
|
|
||||||
_register_success(relay_url)
|
|
||||||
existing_events.extend(events)
|
|
||||||
else:
|
|
||||||
_register_failure(relay_url)
|
|
||||||
|
|
||||||
# Decide whether to publish: publish if none exist or any differ from candidate
|
|
||||||
found_any = len(existing_events) > 0
|
|
||||||
all_match = found_any and all(
|
|
||||||
events_semantically_equal(ev, candidate_event) for ev in existing_events
|
|
||||||
)
|
|
||||||
|
|
||||||
if not all_match:
|
|
||||||
logger.debug(
|
|
||||||
"No matching listing announcement found or differences detected; publishing update"
|
|
||||||
)
|
|
||||||
success_count = 0
|
|
||||||
for relay_url in relay_urls:
|
|
||||||
if _should_skip(relay_url):
|
|
||||||
logger.debug(f"Skipping publish to {relay_url} due to backoff")
|
|
||||||
continue
|
|
||||||
if await publish_to_relay(relay_url, candidate_event):
|
|
||||||
_register_success(relay_url)
|
|
||||||
success_count += 1
|
|
||||||
else:
|
|
||||||
_register_failure(relay_url)
|
|
||||||
logger.info(
|
|
||||||
f"Published listing announcement to {success_count}/{len(relay_urls)} relays"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
logger.debug(
|
|
||||||
"Matching listing announcement already present; skipping publish on startup"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Re-announce periodically (every 24 hours)
|
|
||||||
announcement_interval = 24 * 60 * 60
|
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
await asyncio.sleep(announcement_interval)
|
nsec = (settings.nsec or "").strip()
|
||||||
|
|
||||||
|
if not nsec:
|
||||||
|
if not warned_missing_nsec:
|
||||||
|
logger.info(
|
||||||
|
"Nostr private key not configured (NSEC); waiting for one "
|
||||||
|
"to be set before announcing this provider"
|
||||||
|
)
|
||||||
|
warned_missing_nsec = True
|
||||||
|
parsed_nsec = None
|
||||||
|
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Re-derive the identity whenever the configured NSEC changes, so a
|
||||||
|
# key saved (or replaced) through the admin UI takes effect live.
|
||||||
|
if nsec != parsed_nsec:
|
||||||
|
keypair = nsec_to_keypair(nsec)
|
||||||
|
if not keypair:
|
||||||
|
logger.error(
|
||||||
|
"Invalid NSEC; waiting for a valid one before announcing"
|
||||||
|
)
|
||||||
|
parsed_nsec = None
|
||||||
|
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||||
|
continue
|
||||||
|
private_key_hex, public_key_hex = keypair
|
||||||
|
parsed_nsec = nsec
|
||||||
|
provider_id = None
|
||||||
|
warned_missing_nsec = False
|
||||||
|
logger.info(f"Using Nostr pubkey: {public_key_hex}")
|
||||||
|
|
||||||
|
if private_key_hex is None or public_key_hex is None:
|
||||||
|
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||||
|
continue
|
||||||
|
|
||||||
|
endpoint_urls = _resolve_endpoint_urls()
|
||||||
|
if not endpoint_urls:
|
||||||
|
logger.warning(
|
||||||
|
"No valid endpoints configured (HTTP_URL/ONION_URL). "
|
||||||
|
"Skipping listing publish until one is set."
|
||||||
|
)
|
||||||
|
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||||
|
continue
|
||||||
|
|
||||||
|
relay_urls = _resolve_relay_urls()
|
||||||
|
|
||||||
|
if provider_id is None:
|
||||||
|
provider_id = await _determine_provider_id(public_key_hex, relay_urls)
|
||||||
|
logger.info(f"Using provider_id: {provider_id}")
|
||||||
|
|
||||||
|
metadata = {
|
||||||
|
"name": settings.name or "Routstr Proxy",
|
||||||
|
"about": settings.description
|
||||||
|
or "Privacy-preserving AI proxy via Nostr",
|
||||||
|
}
|
||||||
|
|
||||||
# Build fresh candidate event for comparison
|
|
||||||
version_str = get_app_version()
|
|
||||||
candidate_event = create_listing_event(
|
candidate_event = create_listing_event(
|
||||||
private_key_hex=private_key_hex,
|
private_key_hex=private_key_hex,
|
||||||
provider_id=provider_id,
|
provider_id=provider_id,
|
||||||
endpoint_urls=endpoint_urls,
|
endpoint_urls=endpoint_urls,
|
||||||
mint_urls=mint_urls,
|
mint_urls=_resolve_mint_urls(),
|
||||||
version=version_str,
|
version=get_app_version(),
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Fetch existing events for this provider_id
|
# Fetch existing events for this provider_id
|
||||||
existing_events = []
|
existing_events: list[dict[str, Any]] = []
|
||||||
for relay_url in relay_urls:
|
for relay_url in relay_urls:
|
||||||
if _should_skip(relay_url):
|
if _should_skip(relay_url):
|
||||||
logger.debug(f"Skipping {relay_url} due to backoff")
|
logger.debug(f"Skipping {relay_url} due to backoff")
|
||||||
@@ -549,26 +468,36 @@ async def announce_provider() -> None:
|
|||||||
|
|
||||||
if all_match:
|
if all_match:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Matching listing announcement already present; skipping periodic re-announce"
|
"Matching listing announcement already present; skipping publish"
|
||||||
)
|
)
|
||||||
continue
|
else:
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Re-announcing provider due to differences or absence: {candidate_event['id']}"
|
"No matching listing announcement found or differences "
|
||||||
|
"detected; publishing update"
|
||||||
)
|
)
|
||||||
|
success_count = 0
|
||||||
for relay_url in relay_urls:
|
for relay_url in relay_urls:
|
||||||
if _should_skip(relay_url):
|
if _should_skip(relay_url):
|
||||||
logger.debug(f"Skipping publish to {relay_url} due to backoff")
|
logger.debug(f"Skipping publish to {relay_url} due to backoff")
|
||||||
continue
|
continue
|
||||||
ok = await publish_to_relay(relay_url, candidate_event)
|
if await publish_to_relay(relay_url, candidate_event):
|
||||||
if ok:
|
|
||||||
_register_success(relay_url)
|
_register_success(relay_url)
|
||||||
|
success_count += 1
|
||||||
else:
|
else:
|
||||||
_register_failure(relay_url)
|
_register_failure(relay_url)
|
||||||
|
logger.info(
|
||||||
|
"Published listing announcement to "
|
||||||
|
f"{success_count}/{len(relay_urls)} relays"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Re-announce periodically; wakes early if the NSEC changes.
|
||||||
|
await _sleep_until_next_announcement(
|
||||||
|
ANNOUNCEMENT_INTERVAL_SECONDS, parsed_nsec
|
||||||
|
)
|
||||||
|
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
logger.info("Listing announcement task cancelled")
|
logger.info("Listing announcement task cancelled")
|
||||||
break
|
break
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug(f"Error in listing announcement loop: {type(e).__name__}")
|
logger.debug(f"Error in listing announcement loop: {type(e).__name__}")
|
||||||
# Continue running despite errors
|
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||||
|
|||||||
@@ -0,0 +1,94 @@
|
|||||||
|
"""Small adapter around the maintained ``nostr-sdk`` package."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from datetime import timedelta
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
from nostr_sdk import (
|
||||||
|
AckPolicy,
|
||||||
|
Client,
|
||||||
|
Event,
|
||||||
|
EventBuilder,
|
||||||
|
Filter,
|
||||||
|
Keys,
|
||||||
|
Kind,
|
||||||
|
PublicKey,
|
||||||
|
RelayUrl,
|
||||||
|
ReqExitPolicy,
|
||||||
|
ReqTarget,
|
||||||
|
SendEventTarget,
|
||||||
|
Tag,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_keypair(secret_key: str) -> tuple[str, str]:
|
||||||
|
keys = Keys.parse(secret_key)
|
||||||
|
return keys.secret_key().to_hex(), keys.public_key().to_hex()
|
||||||
|
|
||||||
|
|
||||||
|
def create_signed_event(
|
||||||
|
secret_key_hex: str,
|
||||||
|
*,
|
||||||
|
kind: int,
|
||||||
|
content: str,
|
||||||
|
tags: list[list[str]],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
keys = Keys.parse(secret_key_hex)
|
||||||
|
event = (
|
||||||
|
EventBuilder(Kind(kind), content)
|
||||||
|
.tags([Tag.parse(tag) for tag in tags])
|
||||||
|
.finalize(keys)
|
||||||
|
)
|
||||||
|
return cast(dict[str, Any], json.loads(event.as_json()))
|
||||||
|
|
||||||
|
|
||||||
|
async def fetch_events(
|
||||||
|
relay_url: str,
|
||||||
|
*,
|
||||||
|
kind: int,
|
||||||
|
author: str,
|
||||||
|
limit: int,
|
||||||
|
timeout: int,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
relay = RelayUrl.parse(relay_url)
|
||||||
|
client = Client()
|
||||||
|
await client.add_relay(relay)
|
||||||
|
try:
|
||||||
|
await client.connect()
|
||||||
|
event_filter = (
|
||||||
|
Filter().kinds([Kind(kind)]).authors([PublicKey.parse(author)]).limit(limit)
|
||||||
|
)
|
||||||
|
events = await client.fetch_events(
|
||||||
|
ReqTarget.single(relay, [event_filter]),
|
||||||
|
timeout=timedelta(seconds=timeout),
|
||||||
|
policy=ReqExitPolicy.WAIT_DURATION_AFTER_EOSE(timedelta(seconds=2.5)),
|
||||||
|
max_events=limit,
|
||||||
|
)
|
||||||
|
return [cast(dict[str, Any], json.loads(event.as_json())) for event in events]
|
||||||
|
finally:
|
||||||
|
await client.shutdown()
|
||||||
|
|
||||||
|
|
||||||
|
async def send_event(relay_url: str, event: dict[str, Any], *, timeout: int) -> None:
|
||||||
|
relay = RelayUrl.parse(relay_url)
|
||||||
|
client = Client()
|
||||||
|
await client.add_relay(relay)
|
||||||
|
try:
|
||||||
|
await client.connect()
|
||||||
|
output = await client.send_event(
|
||||||
|
Event.from_json(json.dumps(event)),
|
||||||
|
target=SendEventTarget.to([relay]),
|
||||||
|
ack_policy=AckPolicy.all(),
|
||||||
|
ok_timeout=timedelta(seconds=timeout),
|
||||||
|
)
|
||||||
|
if output.failed or relay not in output.success:
|
||||||
|
reasons = ", ".join(
|
||||||
|
f"{url}: {reason}" for url, reason in output.failed.items()
|
||||||
|
)
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Relay did not accept event: {reasons or 'no OK from relay'}"
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
await client.shutdown()
|
||||||
@@ -1,11 +1,12 @@
|
|||||||
import math
|
import math
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from pydantic.v1 import BaseModel
|
from pydantic.v1 import BaseModel, Field
|
||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
from .price import sats_usd_price
|
from .price import sats_usd_price
|
||||||
|
from .rates import coerce_rate, is_usable_rate
|
||||||
from .usage import normalize_usage, parse_token_count
|
from .usage import normalize_usage, parse_token_count
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -34,6 +35,9 @@ class CostData(BaseModel):
|
|||||||
cache_creation_input_tokens: int = 0
|
cache_creation_input_tokens: int = 0
|
||||||
cache_read_msats: int = 0
|
cache_read_msats: int = 0
|
||||||
cache_creation_msats: int = 0
|
cache_creation_msats: int = 0
|
||||||
|
# Actual debit after finalization; None means settlement has not run yet.
|
||||||
|
charged_msats: int | None = None
|
||||||
|
upstream_usd: float = Field(default=0.0, exclude=True)
|
||||||
|
|
||||||
|
|
||||||
class MaxCostData(CostData):
|
class MaxCostData(CostData):
|
||||||
@@ -107,11 +111,11 @@ async def calculate_cost(
|
|||||||
|
|
||||||
if usage is None:
|
if usage is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"No usage data in response — billing at MaxCostData with zero "
|
"No usage data or local estimate in response — releasing the "
|
||||||
"tokens. Dashboard will show this request as `(0+0)`. Most "
|
"reservation without charging it as usage. Dashboard will show "
|
||||||
"common cause: upstream stream did not include a final usage "
|
"this request as `(0+0)` tokens. Most common cause: upstream "
|
||||||
"chunk (OpenAI-compat backends require "
|
"stream did not include a final usage chunk (OpenAI-compat "
|
||||||
"`stream_options.include_usage=true`).",
|
"backends require `stream_options.include_usage=true`).",
|
||||||
extra={
|
extra={
|
||||||
"max_cost_msats": max_cost,
|
"max_cost_msats": max_cost,
|
||||||
"model": response_data.get("model", "unknown"),
|
"model": response_data.get("model", "unknown"),
|
||||||
@@ -175,13 +179,14 @@ async def calculate_cost(
|
|||||||
cost_details = usage_data.get("cost_details", {})
|
cost_details = usage_data.get("cost_details", {})
|
||||||
if not isinstance(cost_details, dict):
|
if not isinstance(cost_details, dict):
|
||||||
cost_details = {}
|
cost_details = {}
|
||||||
input_usd = _coerce_usd(
|
# Coerce each spelling before choosing between them: `inf` and `NaN`
|
||||||
cost_details.get("input_cost")
|
# are truthy, so a malformed first field would otherwise win the
|
||||||
or cost_details.get("upstream_inference_prompt_cost")
|
# fallback and the usable figure beside it would never be read.
|
||||||
|
input_usd = _coerce_usd(cost_details.get("input_cost")) or _coerce_usd(
|
||||||
|
cost_details.get("upstream_inference_prompt_cost")
|
||||||
)
|
)
|
||||||
output_usd = _coerce_usd(
|
output_usd = _coerce_usd(cost_details.get("output_cost")) or _coerce_usd(
|
||||||
cost_details.get("output_cost")
|
cost_details.get("upstream_inference_completions_cost")
|
||||||
or cost_details.get("upstream_inference_completions_cost")
|
|
||||||
)
|
)
|
||||||
cache_pricing_rates: tuple[float, float, float, float] | None = None
|
cache_pricing_rates: tuple[float, float, float, float] | None = None
|
||||||
if cache_read_tokens > 0 or cache_creation_tokens > 0:
|
if cache_read_tokens > 0 or cache_creation_tokens > 0:
|
||||||
@@ -241,27 +246,32 @@ async def calculate_cost(
|
|||||||
else:
|
else:
|
||||||
input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates
|
input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates
|
||||||
|
|
||||||
if not (input_rate and output_rate):
|
# Truthiness is not the question: `NaN` and a negative rate are both truthy
|
||||||
|
# and sailed past this gate into the token math, while a rate of zero is a
|
||||||
|
# price — free — and reading it as a missing one charged the whole
|
||||||
|
# reservation for a request the model serves for nothing.
|
||||||
|
rates = (input_rate, output_rate, cache_read_rate, cache_creation_rate)
|
||||||
|
if not all(is_usable_rate(rate) for rate in rates):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"No token pricing configured — billing at flat MaxCostData. "
|
"No usable token pricing — releasing the reservation instead of "
|
||||||
"Token counts %s in the upstream response but cannot be "
|
"treating its ceiling as the charge. Token counts %s in the "
|
||||||
"priced; the request will appear in dashboards with the "
|
"upstream response but cannot be converted to money; the request "
|
||||||
"raw counts and a fixed max-cost charge.",
|
"will appear in dashboards with raw counts and a zero charge.",
|
||||||
"are present"
|
"are present" if (input_tokens > 0 or output_tokens > 0) else "are zero",
|
||||||
if (input_tokens > 0 or output_tokens > 0)
|
|
||||||
else "are zero",
|
|
||||||
extra={
|
extra={
|
||||||
"base_cost_msats": max_cost,
|
"base_cost_msats": max_cost,
|
||||||
"model": response_data.get("model", "unknown"),
|
"model": response_data.get("model", "unknown"),
|
||||||
"input_tokens": input_tokens,
|
"input_tokens": input_tokens,
|
||||||
"output_tokens": output_tokens,
|
"output_tokens": output_tokens,
|
||||||
|
"input_rate": input_rate,
|
||||||
|
"output_rate": output_rate,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return MaxCostData(
|
return MaxCostData(
|
||||||
base_msats=max_cost,
|
base_msats=0,
|
||||||
input_msats=0,
|
input_msats=0,
|
||||||
output_msats=0,
|
output_msats=0,
|
||||||
total_msats=max_cost,
|
total_msats=0,
|
||||||
input_tokens=input_tokens,
|
input_tokens=input_tokens,
|
||||||
output_tokens=output_tokens,
|
output_tokens=output_tokens,
|
||||||
cache_read_input_tokens=cache_read_tokens,
|
cache_read_input_tokens=cache_read_tokens,
|
||||||
@@ -289,15 +299,25 @@ async def calculate_cost(
|
|||||||
|
|
||||||
|
|
||||||
def _coerce_usd(value: object) -> float:
|
def _coerce_usd(value: object) -> float:
|
||||||
"""Coerce a value to USD float, handling various formats safely."""
|
"""Coerce an upstream-reported USD figure to a usable amount, else ``0.0``.
|
||||||
if value is None or isinstance(value, bool):
|
|
||||||
return 0.0
|
These values come straight off the upstream response, where ``json.loads``
|
||||||
if not isinstance(value, (int, float, str)):
|
accepts the bare ``NaN``/``Infinity`` literals and overflows ``1e999`` to
|
||||||
return 0.0
|
``inf``. A non-finite figure is not a cost, and letting one through poisoned
|
||||||
try:
|
the proportional split in ``_calculate_from_usd_cost`` (``inf / inf`` is
|
||||||
return max(0.0, float(value))
|
``NaN``): the resulting exception was absorbed by the broad handler around
|
||||||
except (TypeError, ValueError):
|
the USD path, so a request whose *total* cost was perfectly valid fell
|
||||||
return 0.0
|
through to token-estimated pricing and was billed a fraction of what the
|
||||||
|
upstream charged.
|
||||||
|
|
||||||
|
``0.0`` means "no usable figure" to every caller, which is the same thing an
|
||||||
|
absent field means, so the caller's existing ``> 0`` checks handle it.
|
||||||
|
"""
|
||||||
|
# A cost figure is coerced exactly like a rate; only the way an unusable one
|
||||||
|
# is reported differs. A negative is rejected here, where the previous
|
||||||
|
# `max(0.0, …)` clamped it.
|
||||||
|
amount = coerce_rate(value)
|
||||||
|
return amount if amount is not None else 0.0
|
||||||
|
|
||||||
|
|
||||||
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
||||||
@@ -326,9 +346,7 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
|||||||
# actually deducts from the balance. For non-BYOK providers (e.g.
|
# actually deducts from the balance. For non-BYOK providers (e.g.
|
||||||
# OpenRouter) usage.cost already equals upstream_inference_cost, so we
|
# OpenRouter) usage.cost already equals upstream_inference_cost, so we
|
||||||
# fall through to the normal ``cost`` lookup below.
|
# fall through to the normal ``cost`` lookup below.
|
||||||
upstream_cost = _coerce_usd(
|
upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost"))
|
||||||
cost_details.get("upstream_inference_cost")
|
|
||||||
)
|
|
||||||
if upstream_cost > 0 and usage_data.get("is_byok"):
|
if upstream_cost > 0 and usage_data.get("is_byok"):
|
||||||
byok_fee = _coerce_usd(usage_data.get("cost"))
|
byok_fee = _coerce_usd(usage_data.get("cost"))
|
||||||
return upstream_cost + byok_fee
|
return upstream_cost + byok_fee
|
||||||
@@ -359,8 +377,7 @@ def _get_pricing_rates(
|
|||||||
``None`` means configured fixed pricing should be used by the caller.
|
``None`` means configured fixed pricing should be used by the caller.
|
||||||
"""
|
"""
|
||||||
if settings.fixed_pricing and (
|
if settings.fixed_pricing and (
|
||||||
settings.fixed_per_1k_input_tokens
|
settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens
|
||||||
or settings.fixed_per_1k_output_tokens
|
|
||||||
):
|
):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -416,12 +433,8 @@ def _get_pricing_rates(
|
|||||||
usd_per_sat = sats_usd_price()
|
usd_per_sat = sats_usd_price()
|
||||||
mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
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
|
mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||||
cache_read_usd = _coerce_usd(
|
cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost"))
|
||||||
pricing.get("cache_read_input_token_cost")
|
cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost"))
|
||||||
)
|
|
||||||
cache_write_usd = _coerce_usd(
|
|
||||||
pricing.get("cache_creation_input_token_cost")
|
|
||||||
)
|
|
||||||
mscr_1k = (
|
mscr_1k = (
|
||||||
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat
|
||||||
if cache_read_usd > 0
|
if cache_read_usd > 0
|
||||||
@@ -479,6 +492,7 @@ def _calculate_from_usd_cost(
|
|||||||
"""Calculate cost from USD figures, deriving input/output split from tokens."""
|
"""Calculate cost from USD figures, deriving input/output split from tokens."""
|
||||||
if provider_fee is None:
|
if provider_fee is None:
|
||||||
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
|
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
|
||||||
|
reported_usd = usd_cost
|
||||||
usd_cost = usd_cost * provider_fee
|
usd_cost = usd_cost * provider_fee
|
||||||
input_usd = input_usd * provider_fee
|
input_usd = input_usd * provider_fee
|
||||||
output_usd = output_usd * provider_fee
|
output_usd = output_usd * provider_fee
|
||||||
@@ -525,9 +539,7 @@ def _calculate_from_usd_cost(
|
|||||||
regular_weight = input_tokens * input_rate
|
regular_weight = input_tokens * input_rate
|
||||||
cache_read_weight = cache_read_tokens * cache_read_rate
|
cache_read_weight = cache_read_tokens * cache_read_rate
|
||||||
cache_creation_weight = cache_creation_tokens * cache_creation_rate
|
cache_creation_weight = cache_creation_tokens * cache_creation_rate
|
||||||
total_input_weight = (
|
total_input_weight = regular_weight + cache_read_weight + cache_creation_weight
|
||||||
regular_weight + cache_read_weight + cache_creation_weight
|
|
||||||
)
|
|
||||||
if total_input_weight > 0:
|
if total_input_weight > 0:
|
||||||
cache_read_msats = int(
|
cache_read_msats = int(
|
||||||
round(
|
round(
|
||||||
@@ -566,6 +578,7 @@ def _calculate_from_usd_cost(
|
|||||||
cache_creation_input_tokens=cache_creation_tokens,
|
cache_creation_input_tokens=cache_creation_tokens,
|
||||||
cache_read_msats=cache_read_msats,
|
cache_read_msats=cache_read_msats,
|
||||||
cache_creation_msats=cache_creation_msats,
|
cache_creation_msats=cache_creation_msats,
|
||||||
|
upstream_usd=reported_usd,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+285
-52
@@ -1,8 +1,12 @@
|
|||||||
|
import asyncio
|
||||||
import base64
|
import base64
|
||||||
|
import ipaddress
|
||||||
import json
|
import json
|
||||||
import math
|
import math
|
||||||
|
import socket
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
from urllib.parse import urlsplit, urlunsplit
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import HTTPException, Response
|
from fastapi import HTTPException, Response
|
||||||
@@ -14,10 +18,22 @@ from ..core import get_logger
|
|||||||
from ..core.exceptions import UpstreamError
|
from ..core.exceptions import UpstreamError
|
||||||
from ..core.redaction import redact_org_ids
|
from ..core.redaction import redact_org_ids
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
from ..wallet import deserialize_token_from_string
|
from ..wallet import (
|
||||||
|
UntrustedSourceMintError,
|
||||||
|
classify_redemption_error,
|
||||||
|
deserialize_token_from_string,
|
||||||
|
is_trusted_source_mint,
|
||||||
|
)
|
||||||
|
from .responses_input import (
|
||||||
|
FILE_ID_URL_PREFIX,
|
||||||
|
count_input_images,
|
||||||
|
input_image_part_to_image_url,
|
||||||
|
responses_input_to_messages,
|
||||||
|
)
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None:
|
def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None:
|
||||||
if x_cashu := headers.get("x-cashu", None):
|
if x_cashu := headers.get("x-cashu", None):
|
||||||
cashu_token = x_cashu
|
cashu_token = x_cashu
|
||||||
@@ -68,6 +84,19 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
|
|||||||
detail="Invalid authentication token format",
|
detail="Invalid authentication token format",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if not is_trusted_source_mint(token_obj.mint):
|
||||||
|
classified = classify_redemption_error(
|
||||||
|
UntrustedSourceMintError(f"Untrusted source mint: {token_obj.mint}")
|
||||||
|
)
|
||||||
|
assert classified is not None
|
||||||
|
error_type, status_code, message, error_code = classified
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status_code,
|
||||||
|
detail={
|
||||||
|
"error": {"message": message, "type": error_type, "code": error_code}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
amount_msat = (
|
amount_msat = (
|
||||||
token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000
|
token_obj.amount if token_obj.unit == "msat" else token_obj.amount * 1000
|
||||||
)
|
)
|
||||||
@@ -156,7 +185,12 @@ async def calculate_discounted_max_cost(
|
|||||||
body: dict,
|
body: dict,
|
||||||
model_obj: Any | None = None,
|
model_obj: Any | None = None,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""Calculate the discounted max cost for a request using model pricing when available."""
|
"""Calculate the discounted max cost for a request using model pricing when available.
|
||||||
|
|
||||||
|
Completion discounts are trimmed from the largest declared cap among
|
||||||
|
``max_tokens`` and ``max_completion_tokens`` (chat/completions) or
|
||||||
|
``max_output_tokens`` (responses).
|
||||||
|
"""
|
||||||
if settings.fixed_pricing:
|
if settings.fixed_pricing:
|
||||||
return max_cost_for_model
|
return max_cost_for_model
|
||||||
|
|
||||||
@@ -196,10 +230,25 @@ async def calculate_discounted_max_cost(
|
|||||||
|
|
||||||
adjusted = max_cost_for_model
|
adjusted = max_cost_for_model
|
||||||
|
|
||||||
if messages := body.get("messages"):
|
messages = body.get("messages")
|
||||||
prompt_tokens = estimate_tokens(messages)
|
# Estimated over the whole body: a discount driven by message text alone lets
|
||||||
|
# a caller hide prompt weight elsewhere, shrink the reservation, and be billed
|
||||||
|
# for work the reservation never covered.
|
||||||
|
prompt_tokens = estimate_prompt_tokens(body)
|
||||||
|
|
||||||
image_tokens = await estimate_image_tokens_in_messages(messages)
|
# Images are billed as tokens by the upstream but carry no text for
|
||||||
|
# ``estimate_prompt_tokens`` to count, so they are estimated separately and
|
||||||
|
# added on both the chat (``messages``) and Responses (``input``) paths.
|
||||||
|
image_tokens = 0
|
||||||
|
if isinstance(messages, list):
|
||||||
|
image_tokens += await estimate_image_tokens_in_messages(messages)
|
||||||
|
input_data = body.get("input")
|
||||||
|
if isinstance(input_data, list):
|
||||||
|
converted = responses_input_to_messages(input_data)
|
||||||
|
if converted is None:
|
||||||
|
image_tokens += count_input_images(input_data) * _MAX_ORIGINAL_IMAGE_TOKENS
|
||||||
|
else:
|
||||||
|
image_tokens += await estimate_image_tokens_in_messages(converted)
|
||||||
if image_tokens > 0:
|
if image_tokens > 0:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Found images in request",
|
"Found images in request",
|
||||||
@@ -210,22 +259,41 @@ async def calculate_discounted_max_cost(
|
|||||||
)
|
)
|
||||||
prompt_tokens += image_tokens
|
prompt_tokens += image_tokens
|
||||||
|
|
||||||
|
if prompt_tokens > 0:
|
||||||
estimated_prompt_delta_sats = (
|
estimated_prompt_delta_sats = (
|
||||||
max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt
|
max_prompt_allowed_sats - prompt_tokens * model_pricing.prompt
|
||||||
)
|
)
|
||||||
if estimated_prompt_delta_sats > 0:
|
if estimated_prompt_delta_sats > 0:
|
||||||
adjusted = adjusted - math.floor(estimated_prompt_delta_sats * 1000)
|
adjusted = adjusted - math.floor(estimated_prompt_delta_sats * 1000)
|
||||||
|
|
||||||
max_tokens_raw = body.get("max_tokens", None)
|
# Completion caps arrive under several names: ``max_tokens`` (legacy
|
||||||
if max_tokens_raw is not None:
|
# chat), ``max_completion_tokens`` (modern chat) and ``max_output_tokens``
|
||||||
|
# (Responses API). When a request declares more than one, reserve against
|
||||||
|
# the largest: upstream precedence between the fields varies by provider,
|
||||||
|
# so the smaller cap may not be honored and the reservation must never
|
||||||
|
# under-cover what the upstream could bill.
|
||||||
|
max_tokens_int: int | None = None
|
||||||
|
for cap_field in ("max_tokens", "max_completion_tokens", "max_output_tokens"):
|
||||||
|
cap_raw = body.get(cap_field)
|
||||||
|
if cap_raw is None:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
max_tokens_int = int(max_tokens_raw)
|
cap_int = int(cap_raw)
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Invalid max_tokens; ignoring in cost adjustment",
|
"Invalid completion token cap; ignoring in cost adjustment",
|
||||||
extra={"max_tokens": str(max_tokens_raw)[:64], "model": model},
|
extra={
|
||||||
|
"field": cap_field,
|
||||||
|
"value": str(cap_raw)[:64],
|
||||||
|
"model": model,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
else:
|
continue
|
||||||
|
max_tokens_int = (
|
||||||
|
cap_int if max_tokens_int is None else max(max_tokens_int, cap_int)
|
||||||
|
)
|
||||||
|
|
||||||
|
if max_tokens_int is not None:
|
||||||
estimated_completion_delta_sats = (
|
estimated_completion_delta_sats = (
|
||||||
max_completion_allowed_sats - max_tokens_int * model_pricing.completion
|
max_completion_allowed_sats - max_tokens_int * model_pricing.completion
|
||||||
)
|
)
|
||||||
@@ -262,6 +330,51 @@ def estimate_tokens(messages: list) -> int:
|
|||||||
return total // 3
|
return total // 3
|
||||||
|
|
||||||
|
|
||||||
|
def _sum_string_chars(node: Any) -> int:
|
||||||
|
"""Recursively sum the length of every string in the tree, keys included.
|
||||||
|
|
||||||
|
Nothing is excluded. Keys count because JSON-schema property names are
|
||||||
|
forwarded to the provider, and no exclusion rule can be trusted here: every
|
||||||
|
part of the body is caller-controlled, so any carve-out (by key name or by
|
||||||
|
value shape) is a place to hide prompt weight for free. Inline image data is
|
||||||
|
therefore counted as text too, which only makes the discount smaller.
|
||||||
|
"""
|
||||||
|
if isinstance(node, str):
|
||||||
|
return len(node)
|
||||||
|
if isinstance(node, dict):
|
||||||
|
return sum(
|
||||||
|
len(str(key)) + _sum_string_chars(value) for key, value in node.items()
|
||||||
|
)
|
||||||
|
if isinstance(node, list):
|
||||||
|
return sum(_sum_string_chars(item) for item in node)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def _count_prompt_token_ids(node: Any) -> int:
|
||||||
|
if isinstance(node, int) and not isinstance(node, bool):
|
||||||
|
return 1
|
||||||
|
if isinstance(node, list):
|
||||||
|
return sum(_count_prompt_token_ids(item) for item in node)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def estimate_prompt_tokens(body: dict) -> int:
|
||||||
|
"""Conservatively estimate prompt tokens for the whole provider-bound body.
|
||||||
|
|
||||||
|
Every string counts, as do token IDs in legacy ``prompt`` arrays, so no
|
||||||
|
forwarded field can hide prompt weight and shrink its reservation.
|
||||||
|
"""
|
||||||
|
return _sum_string_chars(body) // 3 + _count_prompt_token_ids(body.get("prompt"))
|
||||||
|
|
||||||
|
|
||||||
|
IMAGE_FETCH_TIMEOUT_SECONDS = 10.0
|
||||||
|
# Dimensions live in the header, so a prefix suffices and an endless body cannot
|
||||||
|
# pin memory.
|
||||||
|
IMAGE_FETCH_MAX_BYTES = 512 * 1024
|
||||||
|
# Fetches are sequential, so an unbounded URL list is a request-time amplifier.
|
||||||
|
IMAGE_FETCH_MAX_PER_REQUEST = 8
|
||||||
|
|
||||||
|
|
||||||
def _get_image_dimensions(image_data: bytes) -> tuple[int, int]:
|
def _get_image_dimensions(image_data: bytes) -> tuple[int, int]:
|
||||||
"""Extract image dimensions from image bytes."""
|
"""Extract image dimensions from image bytes."""
|
||||||
try:
|
try:
|
||||||
@@ -275,13 +388,81 @@ def _get_image_dimensions(image_data: bytes) -> tuple[int, int]:
|
|||||||
return (512, 512)
|
return (512, 512)
|
||||||
|
|
||||||
|
|
||||||
async def _fetch_image_from_url(url: str) -> bytes | None:
|
def _is_blocked_address(address: str) -> bool:
|
||||||
"""Fetch image from URL."""
|
"""Allow only globally reachable addresses (RFC 6890)."""
|
||||||
try:
|
try:
|
||||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
ip = ipaddress.ip_address(address)
|
||||||
response = await client.get(url)
|
except ValueError:
|
||||||
|
return True
|
||||||
|
if isinstance(ip, ipaddress.IPv6Address):
|
||||||
|
# An embedded v4 address would otherwise smuggle a rejected target past
|
||||||
|
# the v6 checks.
|
||||||
|
for embedded in (ip.ipv4_mapped, ip.sixtofour):
|
||||||
|
if embedded is not None:
|
||||||
|
return _is_blocked_address(str(embedded))
|
||||||
|
return not ip.is_global or ip.is_multicast
|
||||||
|
|
||||||
|
|
||||||
|
async def _validated_fetch_target(url: str) -> tuple[str, str]:
|
||||||
|
"""Return the URL to request and its ``Host`` header.
|
||||||
|
|
||||||
|
Cost estimation runs on the unauthenticated request body, so a caller can
|
||||||
|
otherwise aim the node at internal hosts. HTTP is rewritten to the resolved
|
||||||
|
address so the name cannot rebind between check and connect; HTTPS keeps its
|
||||||
|
hostname because certificate validation already binds the connection.
|
||||||
|
"""
|
||||||
|
parts = urlsplit(url)
|
||||||
|
if parts.scheme not in ("http", "https"):
|
||||||
|
raise ValueError(f"unsupported scheme: {parts.scheme or 'none'}")
|
||||||
|
host = parts.hostname
|
||||||
|
if not host:
|
||||||
|
raise ValueError("missing host")
|
||||||
|
|
||||||
|
default_port = 443 if parts.scheme == "https" else 80
|
||||||
|
port = parts.port or default_port
|
||||||
|
host_header = f"[{host}]" if ":" in host else host
|
||||||
|
if parts.port is not None:
|
||||||
|
host_header = f"{host_header}:{parts.port}"
|
||||||
|
|
||||||
|
infos = await asyncio.get_running_loop().getaddrinfo(
|
||||||
|
host, port, proto=socket.IPPROTO_TCP
|
||||||
|
)
|
||||||
|
if not infos:
|
||||||
|
raise ValueError("host did not resolve")
|
||||||
|
for info in infos:
|
||||||
|
if _is_blocked_address(str(info[4][0])):
|
||||||
|
raise ValueError("host resolves to a blocked address")
|
||||||
|
|
||||||
|
if parts.scheme == "https":
|
||||||
|
return url, host_header
|
||||||
|
|
||||||
|
family, _, _, _, sockaddr = infos[0]
|
||||||
|
address = str(sockaddr[0])
|
||||||
|
pinned = f"[{address}]" if family == socket.AF_INET6 else address
|
||||||
|
if parts.port is not None:
|
||||||
|
pinned = f"{pinned}:{parts.port}"
|
||||||
|
return urlunsplit((parts.scheme, pinned, parts.path, parts.query, "")), host_header
|
||||||
|
|
||||||
|
|
||||||
|
async def _fetch_image_from_url(url: str) -> bytes | None:
|
||||||
|
"""Fetch the leading bytes of an image, enough to read its dimensions."""
|
||||||
|
try:
|
||||||
|
target, host_header = await _validated_fetch_target(url)
|
||||||
|
async with httpx.AsyncClient(
|
||||||
|
timeout=IMAGE_FETCH_TIMEOUT_SECONDS, follow_redirects=False
|
||||||
|
) as client:
|
||||||
|
async with client.stream(
|
||||||
|
"GET", target, headers={"Host": host_header}
|
||||||
|
) as response:
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
return response.content
|
chunks: list[bytes] = []
|
||||||
|
downloaded = 0
|
||||||
|
async for chunk in response.aiter_bytes():
|
||||||
|
chunks.append(chunk)
|
||||||
|
downloaded += len(chunk)
|
||||||
|
if downloaded >= IMAGE_FETCH_MAX_BYTES:
|
||||||
|
break
|
||||||
|
return b"".join(chunks)[:IMAGE_FETCH_MAX_BYTES]
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Failed to fetch image from URL",
|
"Failed to fetch image from URL",
|
||||||
@@ -290,15 +471,46 @@ async def _fetch_image_from_url(url: str) -> bytes | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
# Patch-based image pricing (OpenAI ``detail: "original"``): the image is
|
||||||
|
# covered with 32x32px patches and billed as ceil(patches * multiplier)
|
||||||
|
# tokens, with no 512px-tile downscaling. The API rejects images above
|
||||||
|
# 30,000 patches, so at the 1.2x multiplier documented for the
|
||||||
|
# original-capable model families (gpt-5.4/5.5/5.6) the worst case a
|
||||||
|
# single image can bill is 36,000 tokens.
|
||||||
|
_IMAGE_PATCH_PX = 32
|
||||||
|
_MAX_IMAGE_PATCHES = 30_000
|
||||||
|
_MAX_ORIGINAL_IMAGE_TOKENS = (_MAX_IMAGE_PATCHES * 6 + 4) // 5 # 36,000
|
||||||
|
|
||||||
|
|
||||||
|
def _calculate_original_image_tokens(width: int, height: int) -> int:
|
||||||
|
"""Estimate tokens for an image billed at ``detail: "original"``.
|
||||||
|
|
||||||
|
Patch-based models cover the image with 32x32px patches and bill
|
||||||
|
``ceil(patches * 1.2)`` tokens. The estimate is bounded by the
|
||||||
|
30,000-patch rejection limit, which is more conservative than the
|
||||||
|
per-model resizing patch budgets (e.g. 10,000 patches on gpt-5.4/5.5)
|
||||||
|
so it never under-reserves.
|
||||||
|
"""
|
||||||
|
patches = ((width + _IMAGE_PATCH_PX - 1) // _IMAGE_PATCH_PX) * (
|
||||||
|
(height + _IMAGE_PATCH_PX - 1) // _IMAGE_PATCH_PX
|
||||||
|
)
|
||||||
|
bounded = min(patches, _MAX_IMAGE_PATCHES)
|
||||||
|
return (bounded * 6 + 4) // 5 # ceil(bounded * 1.2) in exact integer math
|
||||||
|
|
||||||
|
|
||||||
def _calculate_image_tokens(width: int, height: int, detail: str = "auto") -> int:
|
def _calculate_image_tokens(width: int, height: int, detail: str = "auto") -> int:
|
||||||
"""Calculate image tokens based on OpenAI's vision pricing.
|
"""Calculate image tokens based on OpenAI's vision pricing.
|
||||||
|
|
||||||
For low detail: 85 tokens
|
For low detail: 85 tokens
|
||||||
For high detail/auto: 85 base tokens + 170 tokens per 512px tile
|
For high detail/auto: 85 base tokens + 170 tokens per 512px tile
|
||||||
|
For original detail: patch-based pricing at the original resolution
|
||||||
"""
|
"""
|
||||||
if detail == "low":
|
if detail == "low":
|
||||||
return 85
|
return 85
|
||||||
|
|
||||||
|
if detail == "original":
|
||||||
|
return _calculate_original_image_tokens(width, height)
|
||||||
|
|
||||||
if width > 2048 or height > 2048:
|
if width > 2048 or height > 2048:
|
||||||
aspect_ratio = width / height
|
aspect_ratio = width / height
|
||||||
if width > height:
|
if width > height:
|
||||||
@@ -330,6 +542,7 @@ async def estimate_image_tokens_in_messages(messages: list) -> int:
|
|||||||
Supports both base64 encoded images and image URLs.
|
Supports both base64 encoded images and image URLs.
|
||||||
"""
|
"""
|
||||||
total_image_tokens = 0
|
total_image_tokens = 0
|
||||||
|
fetches = 0
|
||||||
|
|
||||||
for message in messages:
|
for message in messages:
|
||||||
if not isinstance(message, dict):
|
if not isinstance(message, dict):
|
||||||
@@ -350,7 +563,9 @@ async def estimate_image_tokens_in_messages(messages: list) -> int:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
content_type = content_item.get("type")
|
content_type = content_item.get("type")
|
||||||
if content_type not in ("image_url", "input_image"):
|
if content_type == "input_image":
|
||||||
|
content_item = input_image_part_to_image_url(content_item)
|
||||||
|
elif content_type != "image_url":
|
||||||
continue
|
continue
|
||||||
|
|
||||||
image_url_data = content_item.get("image_url")
|
image_url_data = content_item.get("image_url")
|
||||||
@@ -362,7 +577,7 @@ async def estimate_image_tokens_in_messages(messages: list) -> int:
|
|||||||
detail = "auto"
|
detail = "auto"
|
||||||
elif isinstance(image_url_data, dict):
|
elif isinstance(image_url_data, dict):
|
||||||
url = image_url_data.get("url", "")
|
url = image_url_data.get("url", "")
|
||||||
detail = image_url_data.get("detail", "auto")
|
detail = image_url_data.get("detail") or "auto"
|
||||||
else:
|
else:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -370,49 +585,67 @@ async def estimate_image_tokens_in_messages(messages: list) -> int:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
if url.startswith("data:image/"):
|
if url.startswith("data:image/"):
|
||||||
try:
|
total_image_tokens += _data_url_image_tokens(url, detail)
|
||||||
header, base64_data = url.split(",", 1)
|
elif url.startswith(FILE_ID_URL_PREFIX):
|
||||||
image_bytes = base64.b64decode(base64_data)
|
total_image_tokens += _worst_case_image_tokens(detail)
|
||||||
width, height = _get_image_dimensions(image_bytes)
|
elif fetches >= IMAGE_FETCH_MAX_PER_REQUEST:
|
||||||
tokens = _calculate_image_tokens(width, height, detail)
|
|
||||||
total_image_tokens += tokens
|
|
||||||
logger.debug(
|
|
||||||
"Calculated tokens for base64 image",
|
|
||||||
extra={
|
|
||||||
"width": width,
|
|
||||||
"height": height,
|
|
||||||
"detail": detail,
|
|
||||||
"tokens": tokens,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Failed to process base64 image",
|
"Skipping image URL fetch above per-request limit",
|
||||||
extra={"error": str(e)},
|
extra={"url": url[:100], "limit": IMAGE_FETCH_MAX_PER_REQUEST},
|
||||||
)
|
)
|
||||||
total_image_tokens += 85
|
total_image_tokens += _worst_case_image_tokens(detail)
|
||||||
else:
|
else:
|
||||||
|
fetches += 1
|
||||||
image_bytes_or_none = await _fetch_image_from_url(url)
|
image_bytes_or_none = await _fetch_image_from_url(url)
|
||||||
if image_bytes_or_none:
|
total_image_tokens += _image_bytes_tokens(
|
||||||
width, height = _get_image_dimensions(image_bytes_or_none)
|
image_bytes_or_none, detail, source=url[:100]
|
||||||
tokens = _calculate_image_tokens(width, height, detail)
|
|
||||||
total_image_tokens += tokens
|
|
||||||
logger.debug(
|
|
||||||
"Calculated tokens for URL image",
|
|
||||||
extra={
|
|
||||||
"url": url[:100],
|
|
||||||
"width": width,
|
|
||||||
"height": height,
|
|
||||||
"detail": detail,
|
|
||||||
"tokens": tokens,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
total_image_tokens += 85
|
|
||||||
|
|
||||||
return total_image_tokens
|
return total_image_tokens
|
||||||
|
|
||||||
|
|
||||||
|
def _worst_case_image_tokens(detail: str) -> int:
|
||||||
|
"""Dimensions unknown: reserve the most ``detail`` can bill."""
|
||||||
|
if detail == "original":
|
||||||
|
return _MAX_ORIGINAL_IMAGE_TOKENS
|
||||||
|
return _calculate_image_tokens(2048, 2048, detail)
|
||||||
|
|
||||||
|
|
||||||
|
def _data_url_image_tokens(url: str, detail: str) -> int:
|
||||||
|
try:
|
||||||
|
_, base64_data = url.split(",", 1)
|
||||||
|
image_bytes = base64.b64decode(base64_data, validate=True)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Failed to decode base64 image", extra={"error": str(e)})
|
||||||
|
return _worst_case_image_tokens(detail)
|
||||||
|
return _image_bytes_tokens(image_bytes, detail, source="data-url")
|
||||||
|
|
||||||
|
|
||||||
|
def _image_bytes_tokens(image_bytes: bytes | None, detail: str, source: str) -> int:
|
||||||
|
if not image_bytes:
|
||||||
|
return _worst_case_image_tokens(detail)
|
||||||
|
try:
|
||||||
|
width, height = Image.open(BytesIO(image_bytes)).size
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to read image dimensions",
|
||||||
|
extra={"error": str(e), "source": source},
|
||||||
|
)
|
||||||
|
return _worst_case_image_tokens(detail)
|
||||||
|
tokens = _calculate_image_tokens(width, height, detail)
|
||||||
|
logger.debug(
|
||||||
|
"Calculated image tokens",
|
||||||
|
extra={
|
||||||
|
"source": source,
|
||||||
|
"width": width,
|
||||||
|
"height": height,
|
||||||
|
"detail": detail,
|
||||||
|
"tokens": tokens,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return tokens
|
||||||
|
|
||||||
|
|
||||||
def create_error_response(
|
def create_error_response(
|
||||||
error_type: str,
|
error_type: str,
|
||||||
message: str,
|
message: str,
|
||||||
|
|||||||
+246
-43
@@ -1,19 +1,28 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import math
|
import asyncio
|
||||||
|
import ipaddress
|
||||||
|
import json
|
||||||
|
import socket
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from typing import TypedDict
|
from typing import Any, TypedDict
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from cashu.core.base import MeltQuoteState
|
from cashu.core.base import MeltQuoteState
|
||||||
|
from cashu.core.settings import settings as cashu_settings
|
||||||
from cashu.wallet.wallet import Proof, Wallet
|
from cashu.wallet.wallet import Proof, Wallet
|
||||||
|
|
||||||
|
from ..cashu_compat import install_cashu_httpx_shim
|
||||||
from ..mint import (
|
from ..mint import (
|
||||||
MINT_TRANSPORT_EXCEPTIONS,
|
|
||||||
is_mint_rate_limited,
|
is_mint_rate_limited,
|
||||||
|
is_mint_transport_error,
|
||||||
run_mint_operation,
|
run_mint_operation,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# cashu 0.20.x passes the `proxies` kwarg httpx removed in 0.28; see the module
|
||||||
|
# docstring. Installed at import so no mint call can run before the patch.
|
||||||
|
install_cashu_httpx_shim()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from bech32 import bech32_decode, convertbits # type: ignore
|
from bech32 import bech32_decode, convertbits # type: ignore
|
||||||
except ModuleNotFoundError: # pragma: no cover – allow runtime miss
|
except ModuleNotFoundError: # pragma: no cover – allow runtime miss
|
||||||
@@ -42,6 +51,114 @@ class MeltOutcomeAmbiguousError(LNURLError):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
class MeltUnpaidError(LNURLError):
|
||||||
|
"""The mint answered the melt request itself with ``unpaid``.
|
||||||
|
|
||||||
|
Unlike :class:`MeltOutcomeAmbiguousError` this is proof that no Lightning
|
||||||
|
payment was made, so callers may restore what they debited.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
_MAX_LNURL_REDIRECTS = 3
|
||||||
|
_MAX_LNURL_RESPONSE_BYTES = 64 * 1024
|
||||||
|
_NON_PUBLIC_HOST_SUFFIXES = (".localhost", ".local", ".internal")
|
||||||
|
|
||||||
|
|
||||||
|
async def _require_public_https_destination(url: httpx.URL) -> None:
|
||||||
|
"""Reject anything that is not a public HTTPS endpoint.
|
||||||
|
|
||||||
|
LNURL destinations and their redirect targets are attacker-influenced, so
|
||||||
|
every hop has to be re-checked: a single ``https://`` origin says nothing
|
||||||
|
about where a 302 points. A bare hostname check is not enough either: a
|
||||||
|
public-looking name can resolve to a loopback/link-local/private address
|
||||||
|
(SSRF), so DNS is resolved here and every resulting address must be global.
|
||||||
|
"""
|
||||||
|
if url.scheme != "https":
|
||||||
|
raise LNURLError("LNURL destination must be an HTTPS URL")
|
||||||
|
|
||||||
|
host = (url.host or "").rstrip(".").lower()
|
||||||
|
if not host:
|
||||||
|
raise LNURLError("LNURL destination has no host")
|
||||||
|
|
||||||
|
try:
|
||||||
|
literal = ipaddress.ip_address(host)
|
||||||
|
except ValueError:
|
||||||
|
literal = None
|
||||||
|
|
||||||
|
if literal is not None:
|
||||||
|
if not literal.is_global:
|
||||||
|
raise LNURLError("LNURL destination is not a public host")
|
||||||
|
return
|
||||||
|
|
||||||
|
if host == "localhost" or host.endswith(_NON_PUBLIC_HOST_SUFFIXES):
|
||||||
|
raise LNURLError("LNURL destination is not a public host")
|
||||||
|
|
||||||
|
port = url.port or 443
|
||||||
|
try:
|
||||||
|
infos = await asyncio.get_running_loop().getaddrinfo(
|
||||||
|
host, port, proto=socket.IPPROTO_TCP
|
||||||
|
)
|
||||||
|
except socket.gaierror as e:
|
||||||
|
raise LNURLError("LNURL destination could not be resolved") from e
|
||||||
|
if not infos:
|
||||||
|
raise LNURLError("LNURL destination could not be resolved")
|
||||||
|
for info in infos:
|
||||||
|
try:
|
||||||
|
resolved = ipaddress.ip_address(info[4][0])
|
||||||
|
except ValueError as e:
|
||||||
|
raise LNURLError("LNURL destination resolved to an invalid address") from e
|
||||||
|
if not resolved.is_global:
|
||||||
|
raise LNURLError("LNURL destination is not a public host")
|
||||||
|
|
||||||
|
|
||||||
|
async def _fetch_lnurl_json(
|
||||||
|
url: str, params: dict[str, int] | None = None
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""GET an LNURL endpoint, validating the destination at every redirect.
|
||||||
|
|
||||||
|
Response bodies are never echoed: an LNURL service is untrusted, and its
|
||||||
|
payload would otherwise reach operator logs through raised errors.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
target = httpx.URL(url, params=params) if params else httpx.URL(url)
|
||||||
|
except httpx.InvalidURL as e:
|
||||||
|
raise LNURLError("LNURL destination is not a usable URL") from e
|
||||||
|
await _require_public_https_destination(target)
|
||||||
|
|
||||||
|
raw: bytes | None = None
|
||||||
|
async with httpx.AsyncClient() as client:
|
||||||
|
for _ in range(_MAX_LNURL_REDIRECTS + 1):
|
||||||
|
async with client.stream(
|
||||||
|
"GET", target, follow_redirects=False, timeout=10
|
||||||
|
) as response:
|
||||||
|
if response.is_redirect:
|
||||||
|
target = target.join(response.headers.get("location", ""))
|
||||||
|
await _require_public_https_destination(target)
|
||||||
|
continue
|
||||||
|
response.raise_for_status()
|
||||||
|
chunks = bytearray()
|
||||||
|
async for chunk in response.aiter_bytes():
|
||||||
|
chunks.extend(chunk)
|
||||||
|
if len(chunks) > _MAX_LNURL_RESPONSE_BYTES:
|
||||||
|
raise LNURLError("LNURL response exceeded the size limit")
|
||||||
|
raw = bytes(chunks)
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
raise LNURLError("LNURL destination exceeded the redirect limit")
|
||||||
|
|
||||||
|
if raw is None:
|
||||||
|
raise LNURLError("LNURL destination exceeded the redirect limit")
|
||||||
|
|
||||||
|
try:
|
||||||
|
data = json.loads(raw)
|
||||||
|
except ValueError as e:
|
||||||
|
raise LNURLError("LNURL response was not valid JSON") from e
|
||||||
|
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
raise LNURLError("LNURL response was not a JSON object")
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
async def decode_lnurl(lnurl: str) -> str:
|
async def decode_lnurl(lnurl: str) -> str:
|
||||||
"""Decode LNURL to get the actual URL.
|
"""Decode LNURL to get the actual URL.
|
||||||
|
|
||||||
@@ -111,26 +228,30 @@ async def get_lnurl_data(lnurl: str) -> LNURLData:
|
|||||||
httpx.HTTPError: If the HTTP request fails
|
httpx.HTTPError: If the HTTP request fails
|
||||||
"""
|
"""
|
||||||
url = await decode_lnurl(lnurl)
|
url = await decode_lnurl(lnurl)
|
||||||
|
lnurl_data = await _fetch_lnurl_json(url)
|
||||||
async with httpx.AsyncClient() as client:
|
|
||||||
response = await client.get(url, follow_redirects=True, timeout=10)
|
|
||||||
response.raise_for_status()
|
|
||||||
|
|
||||||
lnurl_data = response.json()
|
|
||||||
|
|
||||||
# Validate payRequest data
|
# Validate payRequest data
|
||||||
if lnurl_data.get("tag") != "payRequest":
|
if lnurl_data.get("tag") != "payRequest":
|
||||||
raise LNURLError(
|
raise LNURLError("Invalid LNURL tag: expected 'payRequest'")
|
||||||
f"Invalid LNURL tag: expected 'payRequest', got '{lnurl_data.get('tag')}'"
|
|
||||||
)
|
|
||||||
|
|
||||||
if not isinstance(lnurl_data.get("callback"), str):
|
callback_url = lnurl_data.get("callback")
|
||||||
|
if not isinstance(callback_url, str):
|
||||||
raise LNURLError("Invalid LNURL payRequest: missing callback URL")
|
raise LNURLError("Invalid LNURL payRequest: missing callback URL")
|
||||||
|
try:
|
||||||
|
callback_target = httpx.URL(callback_url)
|
||||||
|
except httpx.InvalidURL as e:
|
||||||
|
raise LNURLError("Invalid LNURL callback URL") from e
|
||||||
|
await _require_public_https_destination(callback_target)
|
||||||
|
|
||||||
|
min_sendable = lnurl_data.get("minSendable", 1000) # Default 1 sat
|
||||||
|
max_sendable = lnurl_data.get("maxSendable", 1000000000) # Default 1000 BTC
|
||||||
|
if not isinstance(min_sendable, int) or not isinstance(max_sendable, int):
|
||||||
|
raise LNURLError("Invalid LNURL payRequest: non-integer sendable limits")
|
||||||
|
|
||||||
return LNURLData(
|
return LNURLData(
|
||||||
callback_url=lnurl_data["callback"],
|
callback_url=callback_url,
|
||||||
min_sendable=lnurl_data.get("minSendable", 1000), # Default 1 sat
|
min_sendable=min_sendable,
|
||||||
max_sendable=lnurl_data.get("maxSendable", 1000000000), # Default 1000 BTC
|
max_sendable=max_sendable,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -150,26 +271,53 @@ async def get_lnurl_invoice(
|
|||||||
LNURLError: If the response is invalid
|
LNURLError: If the response is invalid
|
||||||
httpx.HTTPError: If the HTTP request fails
|
httpx.HTTPError: If the HTTP request fails
|
||||||
"""
|
"""
|
||||||
async with httpx.AsyncClient() as client:
|
invoice_data = await _fetch_lnurl_json(callback_url, params={"amount": amount_msat})
|
||||||
response = await client.get(
|
|
||||||
callback_url,
|
|
||||||
params={"amount": amount_msat},
|
|
||||||
follow_redirects=True,
|
|
||||||
timeout=10,
|
|
||||||
)
|
|
||||||
response.raise_for_status()
|
|
||||||
|
|
||||||
invoice_data = response.json()
|
if not isinstance(invoice_data.get("pr"), str):
|
||||||
|
raise LNURLError("LNURL callback returned no invoice")
|
||||||
if "pr" not in invoice_data:
|
|
||||||
# Check if there's an error in the response
|
|
||||||
if "reason" in invoice_data:
|
|
||||||
raise LNURLError(f"LNURL error: {invoice_data['reason']}")
|
|
||||||
raise LNURLError(f"Invalid LNURL invoice response: {invoice_data}")
|
|
||||||
|
|
||||||
return invoice_data["pr"], invoice_data
|
return invoice_data["pr"], invoice_data
|
||||||
|
|
||||||
|
|
||||||
|
def _select_melt_proofs(
|
||||||
|
wallet: Wallet,
|
||||||
|
proofs: list[Proof],
|
||||||
|
*,
|
||||||
|
quote_amount: int,
|
||||||
|
fee_reserve: int,
|
||||||
|
gross_budget: int,
|
||||||
|
) -> tuple[list[Proof] | None, int]:
|
||||||
|
"""Select proofs that cover the quote and exact NUT-02 input fees.
|
||||||
|
|
||||||
|
Cashu 0.20's ``select_to_send`` may recursively swap when asked to spend a
|
||||||
|
wallet's full balance. Melts accept overpayment and return change, so a
|
||||||
|
bounded, largest-first selection is both safer and minimizes input fees.
|
||||||
|
|
||||||
|
Mints reject a melt carrying more than ``mint_max_request_length`` inputs,
|
||||||
|
so a dust-heavy wallet can only pay what its largest inputs cover; the
|
||||||
|
caller lowers the amount and the rest goes out on later payouts.
|
||||||
|
"""
|
||||||
|
selected: list[Proof] = []
|
||||||
|
selected_amount = 0
|
||||||
|
required = quote_amount + fee_reserve
|
||||||
|
spendable = [
|
||||||
|
proof
|
||||||
|
for proof in sorted(proofs, key=lambda item: item.amount, reverse=True)
|
||||||
|
if getattr(proof, "reserved", False) is not True
|
||||||
|
]
|
||||||
|
for proof in spendable[: cashu_settings.mint_max_request_length]:
|
||||||
|
selected.append(proof)
|
||||||
|
selected_amount += proof.amount
|
||||||
|
input_fees = int(wallet.get_fees_for_proofs(selected))
|
||||||
|
required = quote_amount + fee_reserve + input_fees
|
||||||
|
if selected_amount >= required:
|
||||||
|
if required <= gross_budget:
|
||||||
|
return selected, 0
|
||||||
|
# Covered but over budget; more proofs only raise input fees.
|
||||||
|
break
|
||||||
|
return None, max(1, required - min(selected_amount, gross_budget))
|
||||||
|
|
||||||
|
|
||||||
async def raw_send_to_lnurl(
|
async def raw_send_to_lnurl(
|
||||||
wallet: Wallet,
|
wallet: Wallet,
|
||||||
proofs: list[Proof],
|
proofs: list[Proof],
|
||||||
@@ -201,11 +349,10 @@ async def raw_send_to_lnurl(
|
|||||||
# Send USD to Lightning Address
|
# Send USD to Lightning Address
|
||||||
paid = await wallet.send_to_lnurl("user@getalby.com", 50, unit="usd")
|
paid = await wallet.send_to_lnurl("user@getalby.com", 50, unit="usd")
|
||||||
"""
|
"""
|
||||||
total_balance = sum(proof.amount for proof in proofs)
|
if not isinstance(amount, int) or isinstance(amount, bool) or amount <= 0:
|
||||||
if amount and total_balance < amount:
|
raise ValueError("A positive integer amount is required to send to an LNURL.")
|
||||||
|
if sum(proof.amount for proof in proofs) < amount:
|
||||||
raise ValueError("Amount to send is higher than available proofs.")
|
raise ValueError("Amount to send is higher than available proofs.")
|
||||||
else:
|
|
||||||
assert isinstance(amount, int)
|
|
||||||
total_balance = amount
|
total_balance = amount
|
||||||
lnurl_data = await get_lnurl_data(lnurl)
|
lnurl_data = await get_lnurl_data(lnurl)
|
||||||
|
|
||||||
@@ -226,25 +373,51 @@ async def raw_send_to_lnurl(
|
|||||||
f"({min_sendable_sat} - {max_sendable_sat} {unit})"
|
f"({min_sendable_sat} - {max_sendable_sat} {unit})"
|
||||||
)
|
)
|
||||||
|
|
||||||
estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + 1
|
final_amount = amount_msat
|
||||||
estimated_fees_msat = estimated_fees_sat * 1000
|
|
||||||
final_amount = amount_msat - estimated_fees_msat
|
|
||||||
|
|
||||||
|
selected_proofs: list[Proof] | None = None
|
||||||
|
# Find the largest amount covered by the budget after reserve and input fees.
|
||||||
|
for _ in range(8):
|
||||||
|
if final_amount < lnurl_data["min_sendable"]:
|
||||||
|
raise LNURLError("Cashu melt fees leave no payable LNURL amount")
|
||||||
bolt11_invoice, _ = await get_lnurl_invoice(
|
bolt11_invoice, _ = await get_lnurl_invoice(
|
||||||
lnurl_data["callback_url"], final_amount
|
lnurl_data["callback_url"], final_amount
|
||||||
)
|
)
|
||||||
|
|
||||||
melt_quote_resp = await run_mint_operation(
|
melt_quote_resp = await run_mint_operation(
|
||||||
lambda: wallet.melt_quote(invoice=bolt11_invoice),
|
lambda: wallet.melt_quote(invoice=bolt11_invoice),
|
||||||
op_name="lnurl_melt_quote",
|
op_name="lnurl_melt_quote",
|
||||||
mint_url=str(wallet.url),
|
mint_url=str(wallet.url),
|
||||||
|
# Quote creation is unsafe to retry without idempotency.
|
||||||
|
retry_timeouts=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
quoted_amount = int(melt_quote_resp.amount)
|
||||||
|
expected_amount = final_amount // 1000 if unit == "sat" else final_amount
|
||||||
|
if quoted_amount != expected_amount:
|
||||||
|
raise LNURLError(
|
||||||
|
f"LNURL invoice amount does not match the requested amount "
|
||||||
|
f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})"
|
||||||
|
)
|
||||||
|
|
||||||
|
selected_proofs, shortfall = _select_melt_proofs(
|
||||||
|
wallet,
|
||||||
|
proofs,
|
||||||
|
quote_amount=quoted_amount,
|
||||||
|
fee_reserve=int(melt_quote_resp.fee_reserve),
|
||||||
|
gross_budget=amount,
|
||||||
|
)
|
||||||
|
if selected_proofs is not None:
|
||||||
|
break
|
||||||
|
final_amount -= shortfall * (1000 if unit == "sat" else 1)
|
||||||
|
else:
|
||||||
|
raise LNURLError("Cashu melt fees exceed the requested gross amount")
|
||||||
|
|
||||||
if on_melt_quote is not None:
|
if on_melt_quote is not None:
|
||||||
await on_melt_quote(melt_quote_resp.quote)
|
await on_melt_quote(melt_quote_resp.quote)
|
||||||
|
|
||||||
if amount:
|
assert selected_proofs is not None
|
||||||
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
|
proofs = selected_proofs
|
||||||
|
await wallet.set_reserved_for_send(proofs, reserved=True)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
melt_response = await run_mint_operation(
|
melt_response = await run_mint_operation(
|
||||||
@@ -265,15 +438,29 @@ async def raw_send_to_lnurl(
|
|||||||
# reserved as though a Lightning payment could still settle.
|
# reserved as though a Lightning payment could still settle.
|
||||||
await wallet.set_reserved_for_send(proofs, reserved=False)
|
await wallet.set_reserved_for_send(proofs, reserved=False)
|
||||||
raise
|
raise
|
||||||
if not isinstance(error, MINT_TRANSPORT_EXCEPTIONS):
|
if not is_mint_transport_error(error):
|
||||||
raise
|
raise
|
||||||
|
# Cashu clears reservations on transport errors despite an unknown outcome.
|
||||||
|
try:
|
||||||
|
await wallet.set_reserved_for_melt(
|
||||||
|
proofs, reserved=True, quote_id=melt_quote_resp.quote
|
||||||
|
)
|
||||||
|
except Exception as reservation_error:
|
||||||
|
raise MeltOutcomeAmbiguousError(
|
||||||
|
"Melt outcome is ambiguous and its proof reservation could not "
|
||||||
|
"be restored; proofs must not be retried"
|
||||||
|
) from reservation_error
|
||||||
melt_response = None
|
melt_response = None
|
||||||
melt_error: BaseException | None = error
|
melt_error: BaseException | None = error
|
||||||
else:
|
else:
|
||||||
melt_error = None
|
melt_error = None
|
||||||
|
|
||||||
if getattr(melt_response, "state", None) == MeltQuoteState.paid:
|
melt_state = getattr(melt_response, "state", None)
|
||||||
|
if melt_state == MeltQuoteState.paid:
|
||||||
return final_amount
|
return final_amount
|
||||||
|
if melt_state == MeltQuoteState.unpaid:
|
||||||
|
await wallet.set_reserved_for_send(proofs, reserved=False)
|
||||||
|
raise MeltUnpaidError("Cashu mint confirmed that the melt was unpaid")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
quote = await run_mint_operation(
|
quote = await run_mint_operation(
|
||||||
@@ -281,6 +468,8 @@ async def raw_send_to_lnurl(
|
|||||||
op_name="reconcile_lnurl_melt_quote",
|
op_name="reconcile_lnurl_melt_quote",
|
||||||
mint_url=str(wallet.url),
|
mint_url=str(wallet.url),
|
||||||
retry_timeouts=False,
|
retry_timeouts=False,
|
||||||
|
# Reconciliation must bypass the cooldown opened by this failure.
|
||||||
|
allow_during_cooldown=True,
|
||||||
)
|
)
|
||||||
except Exception as reconciliation_error:
|
except Exception as reconciliation_error:
|
||||||
raise MeltOutcomeAmbiguousError(
|
raise MeltOutcomeAmbiguousError(
|
||||||
@@ -290,6 +479,20 @@ async def raw_send_to_lnurl(
|
|||||||
|
|
||||||
if quote is not None and quote.state == MeltQuoteState.paid:
|
if quote is not None and quote.state == MeltQuoteState.paid:
|
||||||
return final_amount
|
return final_amount
|
||||||
|
if quote is not None and quote.state == MeltQuoteState.unpaid:
|
||||||
|
# A just-dispatched quote can briefly report unpaid before transitioning.
|
||||||
|
try:
|
||||||
|
await wallet.set_reserved_for_melt(
|
||||||
|
proofs, reserved=True, quote_id=melt_quote_resp.quote
|
||||||
|
)
|
||||||
|
except Exception as reservation_error:
|
||||||
|
raise MeltOutcomeAmbiguousError(
|
||||||
|
"Melt outcome is ambiguous and its proof reservation could not "
|
||||||
|
"be restored; proofs must not be retried"
|
||||||
|
) from reservation_error
|
||||||
|
raise MeltOutcomeAmbiguousError(
|
||||||
|
"Melt outcome is ambiguous; an immediate unpaid state is not final"
|
||||||
|
) from melt_error
|
||||||
|
|
||||||
state = getattr(getattr(quote, "state", None), "value", "unknown")
|
state = getattr(getattr(quote, "state", None), "value", "unknown")
|
||||||
raise MeltOutcomeAmbiguousError(
|
raise MeltOutcomeAmbiguousError(
|
||||||
|
|||||||
+141
-35
@@ -5,13 +5,14 @@ import random
|
|||||||
import httpx
|
import httpx
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||||
from pydantic import BaseModel as V2BaseModel
|
from pydantic import BaseModel as V2BaseModel
|
||||||
from pydantic.v1 import BaseModel
|
from pydantic.v1 import BaseModel, validator
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from ..core.db import ModelRow, UpstreamProviderRow, get_session
|
from ..core.db import ModelRow, UpstreamProviderRow, get_session
|
||||||
from ..core.logging import get_logger
|
from ..core.logging import get_logger
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
from .price import sats_usd_price
|
from .price import sats_usd_price
|
||||||
|
from .rates import BILLABLE_PRICING_FIELDS, coerce_rate, is_usable_rate
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -58,12 +59,56 @@ class Pricing(BaseModel):
|
|||||||
max_cost: float = 0.0 # in sats not msats
|
max_cost: float = 0.0 # in sats not msats
|
||||||
|
|
||||||
|
|
||||||
|
# The rates ``Pricing`` declares without a default, derived from the model so the
|
||||||
|
# two cannot drift. A payload that omits one writes a row that will not parse.
|
||||||
|
REQUIRED_PRICING_FIELDS = tuple(
|
||||||
|
name for name, field in Pricing.__fields__.items() if field.required
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def has_usable_pricing(pricing: Pricing) -> bool:
|
||||||
|
"""True if every billable rate is a number a request could be billed on.
|
||||||
|
|
||||||
|
Free is usable — a rate of zero is a real price. This asks only whether the
|
||||||
|
price is well-formed. One unusable rate disqualifies the whole price even
|
||||||
|
alongside a valid one, since a request can bill on the bad field: a positive
|
||||||
|
``completion`` must not hide a negative ``prompt``.
|
||||||
|
"""
|
||||||
|
return all(
|
||||||
|
is_usable_rate(getattr(pricing, field)) for field in BILLABLE_PRICING_FIELDS
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TopProvider(BaseModel):
|
class TopProvider(BaseModel):
|
||||||
context_length: int | None = None
|
context_length: int | None = None
|
||||||
max_completion_tokens: int | None = None
|
max_completion_tokens: int | None = None
|
||||||
is_moderated: bool | None = None
|
is_moderated: bool | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class Reasoning(BaseModel):
|
||||||
|
"""Per-model reasoning-effort metadata, matching OpenRouter's shape."""
|
||||||
|
|
||||||
|
mandatory: bool | None = None
|
||||||
|
default_enabled: bool | None = None
|
||||||
|
supported_efforts: list[str] | None = None
|
||||||
|
default_effort: str | None = None
|
||||||
|
supports_max_tokens: bool | None = None
|
||||||
|
|
||||||
|
class Config:
|
||||||
|
extra = "ignore"
|
||||||
|
|
||||||
|
def is_empty(self) -> bool:
|
||||||
|
return not any(
|
||||||
|
(
|
||||||
|
self.mandatory is not None,
|
||||||
|
self.default_enabled is not None,
|
||||||
|
self.supported_efforts,
|
||||||
|
self.default_effort,
|
||||||
|
self.supports_max_tokens is not None,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class Model(BaseModel):
|
class Model(BaseModel):
|
||||||
id: str
|
id: str
|
||||||
name: str
|
name: str
|
||||||
@@ -80,10 +125,43 @@ class Model(BaseModel):
|
|||||||
canonical_slug: str | None = None
|
canonical_slug: str | None = None
|
||||||
alias_ids: list[str] | None = None
|
alias_ids: list[str] | None = None
|
||||||
forwarded_model_id: str | None = None
|
forwarded_model_id: str | None = None
|
||||||
|
reasoning: Reasoning | None = None
|
||||||
|
|
||||||
|
class Config:
|
||||||
|
extra = "ignore"
|
||||||
|
|
||||||
def __hash__(self) -> int:
|
def __hash__(self) -> int:
|
||||||
return hash(self.id)
|
return hash(self.id)
|
||||||
|
|
||||||
|
@validator("reasoning", pre=True)
|
||||||
|
def _coerce_reasoning(cls, value: object) -> object:
|
||||||
|
if value is None or value is False:
|
||||||
|
return None
|
||||||
|
if isinstance(value, Reasoning):
|
||||||
|
return None if value.is_empty() else value
|
||||||
|
if not isinstance(value, dict) or not value:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
parsed = Reasoning.parse_obj(value)
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
return None if parsed.is_empty() else parsed
|
||||||
|
|
||||||
|
def dict(self, **kwargs: object) -> dict:
|
||||||
|
# Non-reasoning models omit the field entirely so the catalog stays
|
||||||
|
# additive: existing clients never see a new null key.
|
||||||
|
data = super().dict(**kwargs) # type: ignore[arg-type]
|
||||||
|
reasoning = data.get("reasoning")
|
||||||
|
if not reasoning:
|
||||||
|
data.pop("reasoning", None)
|
||||||
|
elif isinstance(reasoning, dict):
|
||||||
|
cleaned = {k: v for k, v in reasoning.items() if v is not None}
|
||||||
|
if cleaned:
|
||||||
|
data["reasoning"] = cleaned
|
||||||
|
else:
|
||||||
|
data.pop("reasoning", None)
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
def litellm_cost_entry(model_id: str) -> dict | None:
|
def litellm_cost_entry(model_id: str) -> dict | None:
|
||||||
"""Look up ``model_id`` in litellm's bundled cost map.
|
"""Look up ``model_id`` in litellm's bundled cost map.
|
||||||
@@ -144,18 +222,17 @@ def backfill_cache_pricing(model_id: str, pricing: Pricing) -> Pricing:
|
|||||||
|
|
||||||
|
|
||||||
def _has_valid_pricing(model: dict) -> bool:
|
def _has_valid_pricing(model: dict) -> bool:
|
||||||
"""Check if model has valid pricing (not free, no negative values)."""
|
"""Check if model has valid pricing (usable rates, and not free)."""
|
||||||
pricing = model.get("pricing", {})
|
pricing = model.get("pricing", {})
|
||||||
if not pricing:
|
if not pricing:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
try:
|
# Coercion runs before the both-zero test below, which `NaN` would defeat
|
||||||
prompt = float(pricing.get("prompt", 0))
|
# on its own — and one entry the coercion chokes on must not unwind the
|
||||||
completion = float(pricing.get("completion", 0))
|
# whole fetch, which once cost the node an entire upstream catalog.
|
||||||
except (ValueError, TypeError):
|
prompt = coerce_rate(pricing.get("prompt", 0))
|
||||||
return False
|
completion = coerce_rate(pricing.get("completion", 0))
|
||||||
|
if prompt is None or completion is None:
|
||||||
if prompt < 0 or completion < 0:
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
if prompt == 0 and completion == 0:
|
if prompt == 0 and completion == 0:
|
||||||
@@ -221,9 +298,10 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
||||||
def _row_to_model(
|
def _build_model_from_row(
|
||||||
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
|
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
|
||||||
) -> Model:
|
) -> Model:
|
||||||
|
"""The deterministic USD view of a stored model row, before the sats conversion."""
|
||||||
architecture = json.loads(row.architecture)
|
architecture = json.loads(row.architecture)
|
||||||
pricing = json.loads(row.pricing)
|
pricing = json.loads(row.pricing)
|
||||||
per_request_limits = (
|
per_request_limits = (
|
||||||
@@ -281,6 +359,14 @@ def _row_to_model(
|
|||||||
parsed_pricing.max_cost,
|
parsed_pricing.max_cost,
|
||||||
) = _calculate_usd_max_costs(model)
|
) = _calculate_usd_max_costs(model)
|
||||||
|
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def _row_to_model(
|
||||||
|
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
|
||||||
|
) -> Model:
|
||||||
|
model = _build_model_from_row(row, apply_provider_fee, provider_fee)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
sats_to_usd = sats_usd_price()
|
sats_to_usd = sats_usd_price()
|
||||||
model = _update_model_sats_pricing(model, sats_to_usd)
|
model = _update_model_sats_pricing(model, sats_to_usd)
|
||||||
@@ -309,21 +395,57 @@ async def list_models(
|
|||||||
rows = (await session.exec(query)).all() # type: ignore
|
rows = (await session.exec(query)).all() # type: ignore
|
||||||
provider_result = await session.exec(select(UpstreamProviderRow))
|
provider_result = await session.exec(select(UpstreamProviderRow))
|
||||||
providers_by_id = {p.id: p for p in provider_result.all()}
|
providers_by_id = {p.id: p for p in provider_result.all()}
|
||||||
return [
|
|
||||||
_row_to_model(
|
models: list[Model] = []
|
||||||
|
for r in rows:
|
||||||
|
if not include_disabled and not (
|
||||||
|
r.upstream_provider_id in providers_by_id
|
||||||
|
and providers_by_id[r.upstream_provider_id].enabled
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
model = _row_to_model(
|
||||||
r,
|
r,
|
||||||
apply_provider_fee=apply_fees,
|
apply_provider_fee=apply_fees,
|
||||||
provider_fee=providers_by_id[r.upstream_provider_id].provider_fee
|
provider_fee=providers_by_id[r.upstream_provider_id].provider_fee
|
||||||
if r.upstream_provider_id in providers_by_id
|
if r.upstream_provider_id in providers_by_id
|
||||||
else 1.01,
|
else 1.01,
|
||||||
)
|
)
|
||||||
for r in rows
|
except Exception as e:
|
||||||
if include_disabled
|
# Stored pricing/architecture is JSON from whatever wrote the row, so
|
||||||
or (
|
# a legacy import or foreign writer can leave a field that will not
|
||||||
r.upstream_provider_id in providers_by_id
|
# parse. Converting inside this loop meant one such row raised out of
|
||||||
and providers_by_id[r.upstream_provider_id].enabled
|
# the whole listing and the node advertised nothing at all. Drop the
|
||||||
|
# row we cannot read — it is unservable either way — and keep serving
|
||||||
|
# the rest.
|
||||||
|
logger.warning(
|
||||||
|
"Skipping model row that could not be read",
|
||||||
|
extra={
|
||||||
|
"model_id": r.id,
|
||||||
|
"upstream_provider_id": r.upstream_provider_id,
|
||||||
|
"error": str(e),
|
||||||
|
"error_type": type(e).__name__,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
]
|
continue
|
||||||
|
# Served-map backstop for legacy rows and writers that bypass the admin
|
||||||
|
# edge: a negative or non-finite rate is not a price. Serving one
|
||||||
|
# advertises a rate the cost calculation cannot bill on, so the request
|
||||||
|
# falls through to the flat maximum reservation — or, if the rate is
|
||||||
|
# negative, bills an amount settlement credits back to the caller.
|
||||||
|
# ``include_disabled`` is the operator's listing, which must keep showing
|
||||||
|
# the row so it can be repaired.
|
||||||
|
if not include_disabled and not has_usable_pricing(model.pricing):
|
||||||
|
logger.warning(
|
||||||
|
"Withholding model with an unusable stored rate from the catalog",
|
||||||
|
extra={
|
||||||
|
"model_id": r.id,
|
||||||
|
"upstream_provider_id": r.upstream_provider_id,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
models.append(model)
|
||||||
|
return models
|
||||||
|
|
||||||
|
|
||||||
def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]:
|
def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]:
|
||||||
@@ -409,23 +531,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
|
|||||||
if (sats.max_cost or 0.0) < min_req_sats:
|
if (sats.max_cost or 0.0) < min_req_sats:
|
||||||
sats.max_cost = min_req_sats
|
sats.max_cost = min_req_sats
|
||||||
|
|
||||||
return Model(
|
return model.copy(update={"sats_pricing": sats})
|
||||||
id=model.id,
|
|
||||||
name=model.name,
|
|
||||||
created=model.created,
|
|
||||||
description=model.description,
|
|
||||||
context_length=model.context_length,
|
|
||||||
architecture=model.architecture,
|
|
||||||
pricing=model.pricing,
|
|
||||||
sats_pricing=sats,
|
|
||||||
per_request_limits=model.per_request_limits,
|
|
||||||
top_provider=model.top_provider,
|
|
||||||
enabled=model.enabled,
|
|
||||||
upstream_provider_id=model.upstream_provider_id,
|
|
||||||
canonical_slug=model.canonical_slug,
|
|
||||||
alias_ids=model.alias_ids,
|
|
||||||
forwarded_model_id=model.forwarded_model_id,
|
|
||||||
)
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(
|
logger.error(
|
||||||
"Failed to update sats pricing for model",
|
"Failed to update sats pricing for model",
|
||||||
|
|||||||
+35
-13
@@ -5,12 +5,37 @@ import httpx
|
|||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
|
from .rates import coerce_rate
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
BTC_USD_PRICE: float | None = None
|
BTC_USD_PRICE: float | None = None
|
||||||
SATS_USD_PRICE: float | None = None
|
SATS_USD_PRICE: float | None = None
|
||||||
|
|
||||||
|
SATS_PER_BTC = 100_000_000
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_quote(raw: object, exchange: str) -> float | None:
|
||||||
|
"""Coerce an exchange quote to a price, or ``None`` if it is not one.
|
||||||
|
|
||||||
|
Every quote passes through here because the aggregator takes the ``min()``
|
||||||
|
of what it collects: an unusable quote does not merely join the sample, it
|
||||||
|
*wins* it, and the result is the rate every model and every request on the
|
||||||
|
node is priced at. A quote is stricter than a billable rate — it must be
|
||||||
|
positive, and positive *after* the sats conversion the node prices in: a
|
||||||
|
subnormal quote survives every guard here and still underflows to a zero
|
||||||
|
sats price, which then divides by zero on every model's rate.
|
||||||
|
"""
|
||||||
|
price = coerce_rate(raw)
|
||||||
|
if price is None or price <= 0 or price / SATS_PER_BTC <= 0:
|
||||||
|
logger.warning(
|
||||||
|
"Unusable price quote — ignoring this exchange",
|
||||||
|
extra={"exchange": exchange, "quote": repr(raw)},
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
return price
|
||||||
|
|
||||||
|
|
||||||
async def _kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
|
async def _kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
|
||||||
"""Fetch BTC/USD price from Kraken API."""
|
"""Fetch BTC/USD price from Kraken API."""
|
||||||
@@ -18,10 +43,11 @@ async def _kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
|
|||||||
try:
|
try:
|
||||||
response = await client.get(api)
|
response = await client.get(api)
|
||||||
price_data = response.json()
|
price_data = response.json()
|
||||||
price = float(price_data["result"]["XXBTZUSD"]["c"][0])
|
return _parse_quote(price_data["result"]["XXBTZUSD"]["c"][0], "kraken")
|
||||||
|
except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e:
|
||||||
return price
|
# A payload whose *shape* changed raises IndexError/TypeError, and a
|
||||||
except (httpx.RequestError, KeyError) as e:
|
# non-JSON body raises ValueError; unhandled, one exchange's bad day
|
||||||
|
# aborted the whole aggregation instead of dropping a single quote.
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Kraken API error",
|
"Kraken API error",
|
||||||
extra={
|
extra={
|
||||||
@@ -39,10 +65,8 @@ async def _coinbase_btc_usd(client: httpx.AsyncClient) -> float | None:
|
|||||||
try:
|
try:
|
||||||
response = await client.get(api)
|
response = await client.get(api)
|
||||||
price_data = response.json()
|
price_data = response.json()
|
||||||
price = float(price_data["data"]["amount"])
|
return _parse_quote(price_data["data"]["amount"], "coinbase")
|
||||||
|
except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e:
|
||||||
return price
|
|
||||||
except (httpx.RequestError, KeyError) as e:
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Coinbase API error",
|
"Coinbase API error",
|
||||||
extra={
|
extra={
|
||||||
@@ -60,10 +84,8 @@ async def _binance_btc_usdt(client: httpx.AsyncClient) -> float | None:
|
|||||||
try:
|
try:
|
||||||
response = await client.get(api)
|
response = await client.get(api)
|
||||||
price_data = response.json()
|
price_data = response.json()
|
||||||
price = float(price_data["price"])
|
return _parse_quote(price_data["price"], "binance")
|
||||||
|
except (httpx.RequestError, KeyError, IndexError, TypeError, ValueError) as e:
|
||||||
return price
|
|
||||||
except (httpx.RequestError, KeyError) as e:
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Binance API error",
|
"Binance API error",
|
||||||
extra={
|
extra={
|
||||||
@@ -123,7 +145,7 @@ async def _update_prices() -> None:
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
BTC_USD_PRICE = btc_price
|
BTC_USD_PRICE = btc_price
|
||||||
SATS_USD_PRICE = btc_price / 100_000_000
|
SATS_USD_PRICE = btc_price / SATS_PER_BTC
|
||||||
|
|
||||||
|
|
||||||
def btc_usd_price() -> float:
|
def btc_usd_price() -> float:
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
"""The one definition of a billable rate, with no dependencies of its own.
|
||||||
|
|
||||||
|
A rate reaches the node from an upstream catalog, the LiteLLM cost map, an
|
||||||
|
operator's admin edit, a legacy database row and the BTC/USD feed. Each of those
|
||||||
|
readers needs the same two questions answered — is this value a rate at all, and
|
||||||
|
is it a rate a request can be billed on — so both answers live here, in a module
|
||||||
|
that imports nothing from the package. Every guard then shares one definition
|
||||||
|
instead of drifting, and no caller needs a deferred import to reach it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import math
|
||||||
|
|
||||||
|
# The rates a request can bill on. Derived fields (``max_*_cost``) are excluded
|
||||||
|
# — they are computed carriers, not charged rates. One definition, shared by the
|
||||||
|
# admin write edge and the served/routed guards, so they all cover the same set.
|
||||||
|
BILLABLE_PRICING_FIELDS = (
|
||||||
|
"prompt",
|
||||||
|
"completion",
|
||||||
|
"request",
|
||||||
|
"image",
|
||||||
|
"web_search",
|
||||||
|
"internal_reasoning",
|
||||||
|
"input_cache_read",
|
||||||
|
"input_cache_write",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_usable_rate(rate: float) -> bool:
|
||||||
|
"""True if a single billable rate is a number a request could be billed on.
|
||||||
|
|
||||||
|
The one definition of a usable rate, so every guard that asks the question
|
||||||
|
answers it identically. A rate qualifies only when it is finite and
|
||||||
|
non-negative; zero is usable (it means "free", which is a real price) but
|
||||||
|
``NaN``, ``±inf`` and negatives are not prices at all.
|
||||||
|
|
||||||
|
Non-finite: ``inf > 0`` is True, so an infinite rate reads as chargeable and
|
||||||
|
would be served, routed and billed as ``inf``; ``NaN`` poisons every total it
|
||||||
|
enters and defeats ordinary comparisons, since ``NaN > 0``, ``NaN < 0`` and
|
||||||
|
``NaN == 0`` are all False. Negative: a negative rate produces a negative
|
||||||
|
cost, which the settlement path subtracts from the balance — it pays the
|
||||||
|
caller to make requests. Both reach a stored row from upstream catalogs as
|
||||||
|
well as the admin edge (``json.loads`` accepts the bare ``NaN``/``Infinity``
|
||||||
|
literals and overflows ``1e999`` to ``inf``).
|
||||||
|
|
||||||
|
This is the rationale for every guard that calls it; the call sites say what
|
||||||
|
they do with the answer, not why the answer matters.
|
||||||
|
"""
|
||||||
|
return math.isfinite(rate) and rate >= 0.0
|
||||||
|
|
||||||
|
|
||||||
|
def coerce_rate(value: object) -> float | None:
|
||||||
|
"""Coerce a value from outside the node to a usable rate, or ``None``.
|
||||||
|
|
||||||
|
The one coercion, shared by every reader of a rate the node did not compute
|
||||||
|
itself: an upstream catalog, the LiteLLM cost map, the exchange feed and the
|
||||||
|
admin write edge. Each of them was parsing for itself, and they disagreed —
|
||||||
|
which is how a boolean became a price on some paths and not others.
|
||||||
|
|
||||||
|
A boolean is rejected outright: it is a change of shape, not a rate, and
|
||||||
|
Python would make ``True`` a finite, positive ``1.0`` that passes every
|
||||||
|
numeric guard downstream — a dollar per token. A numeric string is accepted,
|
||||||
|
because feeds report prices as strings. An oversized integer raises
|
||||||
|
``OverflowError`` rather than ``ValueError``, so that is caught too.
|
||||||
|
"""
|
||||||
|
if isinstance(value, bool) or not isinstance(value, (int, float, str)):
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
rate = float(value)
|
||||||
|
except (TypeError, ValueError, OverflowError):
|
||||||
|
return None
|
||||||
|
return rate if is_usable_rate(rate) else None
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
"""Convert a Responses API ``input`` into chat ``messages`` via litellm.
|
||||||
|
|
||||||
|
litellm drops ``file_id`` (emits ``url: ""``) and nests a dict-form ``image_url``
|
||||||
|
as-is, so ``input_image`` parts are flattened to ``{image_url: str, detail}`` first.
|
||||||
|
``file_id`` becomes a sentinel URL the image walker treats as unfetchable.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||||
|
LiteLLMCompletionResponsesConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
from ..core import get_logger
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
FILE_ID_URL_PREFIX = "file-id:"
|
||||||
|
|
||||||
|
|
||||||
|
def _flatten_input_image(part: dict[str, Any]) -> tuple[str, str]:
|
||||||
|
raw = part.get("image_url")
|
||||||
|
url = raw.get("url", "") if isinstance(raw, dict) else raw
|
||||||
|
detail = part.get("detail") or (
|
||||||
|
raw.get("detail") if isinstance(raw, dict) else None
|
||||||
|
)
|
||||||
|
if not url and part.get("file_id"):
|
||||||
|
url = f"{FILE_ID_URL_PREFIX}{part['file_id']}"
|
||||||
|
return (url if isinstance(url, str) else ""), (detail or "auto")
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_item(item: Any) -> Any:
|
||||||
|
if not isinstance(item, dict):
|
||||||
|
return item
|
||||||
|
if item.get("type") == "input_image":
|
||||||
|
url, detail = _flatten_input_image(item)
|
||||||
|
return {**item, "image_url": url, "detail": detail}
|
||||||
|
content = item.get("content")
|
||||||
|
if isinstance(content, list):
|
||||||
|
return {**item, "content": [_normalize_item(part) for part in content]}
|
||||||
|
return item
|
||||||
|
|
||||||
|
|
||||||
|
def input_image_part_to_image_url(part: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
"""Reshape an ``input_image`` part found inside chat ``messages``."""
|
||||||
|
url, detail = _flatten_input_image(part)
|
||||||
|
return {"type": "image_url", "image_url": {"url": url, "detail": detail}}
|
||||||
|
|
||||||
|
|
||||||
|
def count_input_images(input_data: Any) -> int:
|
||||||
|
if isinstance(input_data, dict):
|
||||||
|
own = 1 if input_data.get("type") == "input_image" else 0
|
||||||
|
return own + count_input_images(input_data.get("content"))
|
||||||
|
if isinstance(input_data, list):
|
||||||
|
return sum(count_input_images(item) for item in input_data)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
def responses_input_to_messages(input_data: Any) -> list[dict[str, Any]] | None:
|
||||||
|
"""Returns ``None`` when the transform fails so the caller can worst-case."""
|
||||||
|
if isinstance(input_data, str):
|
||||||
|
return [{"role": "user", "content": input_data}]
|
||||||
|
if not isinstance(input_data, list):
|
||||||
|
return []
|
||||||
|
try:
|
||||||
|
normalized = [_normalize_item(item) for item in input_data]
|
||||||
|
converted = (
|
||||||
|
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||||
|
input=normalized, # type: ignore[arg-type]
|
||||||
|
responses_api_request={},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return [dict(message) for message in converted]
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"Responses input transform failed; using conservative image fallback",
|
||||||
|
extra={"error": str(e)},
|
||||||
|
)
|
||||||
|
return None
|
||||||
+276
-26
@@ -26,6 +26,7 @@ from .core.db import (
|
|||||||
)
|
)
|
||||||
from .core.exceptions import UpstreamError
|
from .core.exceptions import UpstreamError
|
||||||
from .core.not_found import build_not_found_response
|
from .core.not_found import build_not_found_response
|
||||||
|
from .core.settings import settings
|
||||||
from .payment.helpers import (
|
from .payment.helpers import (
|
||||||
calculate_discounted_max_cost,
|
calculate_discounted_max_cost,
|
||||||
check_token_balance,
|
check_token_balance,
|
||||||
@@ -37,9 +38,18 @@ from .payment.models import Model
|
|||||||
from .upstream import BaseUpstreamProvider
|
from .upstream import BaseUpstreamProvider
|
||||||
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
|
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
|
||||||
from .upstream.helpers import init_upstreams
|
from .upstream.helpers import init_upstreams
|
||||||
|
from .upstream.model_paths import (
|
||||||
|
ModelPathSelector,
|
||||||
|
decode_model_path,
|
||||||
|
is_openrouter_base_url,
|
||||||
|
public_model_id,
|
||||||
|
public_provider_url,
|
||||||
|
)
|
||||||
from .upstream.request_correction import correct_request, extract_error_message
|
from .upstream.request_correction import correct_request, extract_error_message
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
MODEL_PATH_HEADER = "x-routstr-model-path"
|
||||||
proxy_router = APIRouter()
|
proxy_router = APIRouter()
|
||||||
|
|
||||||
_upstreams: list[BaseUpstreamProvider] = []
|
_upstreams: list[BaseUpstreamProvider] = []
|
||||||
@@ -112,6 +122,25 @@ def get_candidates(
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _model_ids_match(requested: str, selected: str) -> bool:
|
||||||
|
if requested.lower() == selected.lower():
|
||||||
|
return True
|
||||||
|
return public_model_id(requested).lower() == public_model_id(selected).lower()
|
||||||
|
|
||||||
|
|
||||||
|
def _candidate_for_selector(
|
||||||
|
selector: ModelPathSelector,
|
||||||
|
candidates: list[tuple[Model, BaseUpstreamProvider]],
|
||||||
|
) -> tuple[Model, BaseUpstreamProvider] | None:
|
||||||
|
for model_obj, upstream in candidates:
|
||||||
|
if (
|
||||||
|
upstream.db_id == selector.provider_id
|
||||||
|
and public_provider_url(upstream.base_url) == selector.base_url
|
||||||
|
):
|
||||||
|
return model_obj, upstream
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def get_model_instance(model_id: str) -> Model | None:
|
def get_model_instance(model_id: str) -> Model | None:
|
||||||
"""Get the best-ranked Model instance for a model ID."""
|
"""Get the best-ranked Model instance for a model ID."""
|
||||||
candidates = get_candidates(model_id)
|
candidates = get_candidates(model_id)
|
||||||
@@ -221,22 +250,141 @@ async def refresh_model_maps_periodically() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
_API_PATH_PREFIXES = (
|
# Canonical endpoints this proxy will forward, keyed by the path with any
|
||||||
"v1/",
|
# leading "v1/" and trailing slash removed, mapped to the methods allowed on
|
||||||
"responses",
|
# each. The provider credential is attached during forwarding, so endpoint
|
||||||
"chat/",
|
# permission has to come from this table rather than from the client-supplied
|
||||||
"completions",
|
# path: an upstream's key-management, organization, or billing routes live
|
||||||
"models",
|
# under the same origin and must never be reachable through the proxy.
|
||||||
"embeddings",
|
_ALLOWED_ENDPOINTS: dict[str, frozenset[str]] = {
|
||||||
"audio/",
|
"chat/completions": frozenset({"POST"}),
|
||||||
"images/",
|
"completions": frozenset({"POST"}),
|
||||||
"moderations",
|
"responses": frozenset({"POST"}),
|
||||||
"providers",
|
"messages": frozenset({"POST"}),
|
||||||
"tee/",
|
"embeddings": frozenset({"POST"}),
|
||||||
"attestation",
|
"models": frozenset({"GET"}),
|
||||||
|
"attestation": frozenset({"GET"}),
|
||||||
|
"tee/attestation": frozenset({"GET"}),
|
||||||
|
}
|
||||||
|
|
||||||
|
_ALLOWED_METHODS = frozenset({"GET", "POST"})
|
||||||
|
|
||||||
|
|
||||||
|
def _canonical_api_path(path: str) -> str:
|
||||||
|
"""Reduce a request path to its allowlist key.
|
||||||
|
|
||||||
|
OpenAI-style clients reach the same endpoint with or without the ``v1/``
|
||||||
|
prefix and with or without a trailing slash, so both spellings collapse to
|
||||||
|
one key. Callers must screen the path with
|
||||||
|
:func:`_is_ambiguously_spelled_path` first — this function assumes the path
|
||||||
|
has no dot segments, empty segments, or encoded separators left to resolve.
|
||||||
|
"""
|
||||||
|
core = path[:-1] if path.endswith("/") else path
|
||||||
|
if core.startswith("v1/"):
|
||||||
|
core = core[len("v1/") :]
|
||||||
|
return core
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_extra_allowed_endpoints(raw: str) -> dict[str, frozenset[str]]:
|
||||||
|
"""Parse operator-configured additions to the endpoint allowlist.
|
||||||
|
|
||||||
|
Deployments whose provider exposes an endpoint outside the canonical set
|
||||||
|
opt in explicitly with ``PROXY_EXTRA_ALLOWED_PATHS``, a comma-separated
|
||||||
|
list of ``METHOD:path`` pairs (e.g. ``POST:v1/rerank,GET:batches``). Every
|
||||||
|
entry must name one concrete method and one unambiguous path; wildcards
|
||||||
|
and bare prefixes are deliberately unsupported, so widening the proxy's
|
||||||
|
reach is always a per-endpoint decision. Malformed entries are dropped
|
||||||
|
with a warning rather than silently widening or narrowing the surface.
|
||||||
|
"""
|
||||||
|
extra: dict[str, frozenset[str]] = {}
|
||||||
|
for entry in raw.split(","):
|
||||||
|
entry = entry.strip()
|
||||||
|
if not entry:
|
||||||
|
continue
|
||||||
|
method, separator, endpoint = entry.partition(":")
|
||||||
|
method = method.strip().upper()
|
||||||
|
endpoint = endpoint.strip()
|
||||||
|
if not separator or method not in _ALLOWED_METHODS or not endpoint:
|
||||||
|
logger.warning(
|
||||||
|
"Ignoring malformed PROXY_EXTRA_ALLOWED_PATHS entry",
|
||||||
|
extra={"entry": entry},
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
if _is_ambiguously_spelled_path(endpoint):
|
||||||
|
logger.warning(
|
||||||
|
"Ignoring ambiguously spelled PROXY_EXTRA_ALLOWED_PATHS entry",
|
||||||
|
extra={"entry": entry},
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
if any(character in endpoint for character in "*?["):
|
||||||
|
# Refuse glob syntax outright. Kept as a literal endpoint name it
|
||||||
|
# would never match a real request, so the operator would think
|
||||||
|
# they had widened the proxy when they had not.
|
||||||
|
logger.warning(
|
||||||
|
"Ignoring wildcard PROXY_EXTRA_ALLOWED_PATHS entry; "
|
||||||
|
"list each endpoint explicitly",
|
||||||
|
extra={"entry": entry},
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
key = _canonical_api_path(endpoint)
|
||||||
|
extra[key] = extra.get(key, frozenset()) | {method}
|
||||||
|
return extra
|
||||||
|
|
||||||
|
|
||||||
|
def _is_ambiguously_spelled_path(path: str) -> bool:
|
||||||
|
"""Reject paths whose spelling could resolve somewhere the allowlist did not.
|
||||||
|
|
||||||
|
``{path:path}`` arrives percent-decoded, so a client that sent ``%2e%2e`` or
|
||||||
|
``%2f`` shows up here as ``..`` / ``/``. Dot segments, backslashes, duplicate
|
||||||
|
or leading separators, NUL bytes, and any residual encoded separator are
|
||||||
|
treated as unsafe: they let a caller walk off the canonical API surface (and
|
||||||
|
onto a sensitive upstream endpoint) even though the literal prefix check
|
||||||
|
would pass. Reject rather than trying to rewrite the path.
|
||||||
|
"""
|
||||||
|
if not path or path != path.strip() or path.startswith("/"):
|
||||||
|
return True
|
||||||
|
if "\x00" in path or "\\" in path:
|
||||||
|
return True
|
||||||
|
# A single trailing slash is canonical (e.g. "attestation/"); ignore it,
|
||||||
|
# then no remaining segment may be empty (covers "//") or a dot segment.
|
||||||
|
core = path[:-1] if path.endswith("/") else path
|
||||||
|
if any(segment in ("", ".", "..") for segment in core.split("/")):
|
||||||
|
return True
|
||||||
|
lowered = path.lower()
|
||||||
|
return "%2e" in lowered or "%2f" in lowered or "%5c" in lowered
|
||||||
|
|
||||||
|
|
||||||
|
_EXTRA_ALLOWED_ENDPOINTS = _parse_extra_allowed_endpoints(
|
||||||
|
settings.proxy_extra_allowed_paths
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _allowed_methods_for(endpoint: str) -> frozenset[str]:
|
||||||
|
"""Return the methods allowed on a canonical endpoint, empty if unknown."""
|
||||||
|
methods = _ALLOWED_ENDPOINTS.get(endpoint, frozenset())
|
||||||
|
methods |= _EXTRA_ALLOWED_ENDPOINTS.get(endpoint, frozenset())
|
||||||
|
return methods
|
||||||
|
|
||||||
|
|
||||||
|
def _forwarding_allowed(path: str, method: str) -> bool:
|
||||||
|
"""Gate which method/path pairs may reach an upstream at all.
|
||||||
|
|
||||||
|
The provider credential is attached during forwarding, so an unknown
|
||||||
|
endpoint must never be forwarded on the caller's say-so. The path is
|
||||||
|
reduced to its canonical form and looked up in the endpoint table; there is
|
||||||
|
no prefix match, so a known prefix no longer carries an unknown endpoint
|
||||||
|
(``v1/organization/api_keys`` is rejected even though ``v1/`` is familiar).
|
||||||
|
|
||||||
|
EHBP requests are gated by the same table. Their body is opaque to the
|
||||||
|
proxy, which is a reason to constrain the destination more tightly, not to
|
||||||
|
trust the caller's path: the encrypted contract covers the body, never the
|
||||||
|
endpoint the credential is spent against.
|
||||||
|
"""
|
||||||
|
if method not in _ALLOWED_METHODS:
|
||||||
|
return False
|
||||||
|
return method in _allowed_methods_for(_canonical_api_path(path))
|
||||||
|
|
||||||
|
|
||||||
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
|
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
|
||||||
async def proxy(
|
async def proxy(
|
||||||
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
||||||
@@ -255,14 +403,17 @@ async def proxy(
|
|||||||
async def _proxy(
|
async def _proxy(
|
||||||
request: Request, path: str, session: AsyncSession
|
request: Request, path: str, session: AsyncSession
|
||||||
) -> Response | StreamingResponse:
|
) -> Response | StreamingResponse:
|
||||||
# GET requests must hit a known API prefix; otherwise return a 404 (HTML
|
# Screen the path before any routing decision: reject ambiguous spellings,
|
||||||
# for browsers, JSON for API clients). POST requests are always forwarded
|
# then require a known API prefix so nothing unknown is forwarded with the
|
||||||
# so that OpenAI-style endpoints work with or without the `v1/` prefix
|
# provider credential attached.
|
||||||
# (e.g. `/chat/completions` as well as `/v1/chat/completions`).
|
if _is_ambiguously_spelled_path(path):
|
||||||
if request.method == "GET" and not path.startswith(_API_PATH_PREFIXES):
|
|
||||||
return build_not_found_response(request, path)
|
return build_not_found_response(request, path)
|
||||||
|
|
||||||
headers = dict(request.headers)
|
headers = dict(request.headers)
|
||||||
|
is_ehbp = "ehbp-encapsulated-key" in headers
|
||||||
|
|
||||||
|
if not _forwarding_allowed(path, request.method):
|
||||||
|
return build_not_found_response(request, path)
|
||||||
|
|
||||||
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
|
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
|
||||||
request_body = await request.body()
|
request_body = await request.body()
|
||||||
@@ -272,7 +423,6 @@ async def _proxy(
|
|||||||
# extract the model id, so the SDK sends it in X-Routstr-Model. Forward the
|
# extract the model id, so the SDK sends it in X-Routstr-Model. Forward the
|
||||||
# raw encrypted body to the upstream's /private/ endpoint and stream the
|
# raw encrypted body to the upstream's /private/ endpoint and stream the
|
||||||
# encrypted response back untouched — the SDK's SecureClient decrypts it.
|
# encrypted response back untouched — the SDK's SecureClient decrypts it.
|
||||||
is_ehbp = "ehbp-encapsulated-key" in headers
|
|
||||||
if is_ehbp:
|
if is_ehbp:
|
||||||
request_body_dict = {}
|
request_body_dict = {}
|
||||||
model_id = headers.get("x-routstr-model", "")
|
model_id = headers.get("x-routstr-model", "")
|
||||||
@@ -294,6 +444,13 @@ async def _proxy(
|
|||||||
# without model/cost/auth lookups. Do not prefix-match here: paths such as
|
# without model/cost/auth lookups. Do not prefix-match here: paths such as
|
||||||
# /attestationjunk must continue through normal authentication.
|
# /attestationjunk must continue through normal authentication.
|
||||||
if request.method == "GET" and _is_tinfoil_attestation_path(path):
|
if request.method == "GET" and _is_tinfoil_attestation_path(path):
|
||||||
|
if MODEL_PATH_HEADER in headers:
|
||||||
|
return create_error_response(
|
||||||
|
"unsupported_request",
|
||||||
|
"Model paths do not apply to attestation",
|
||||||
|
400,
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams)
|
selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams)
|
||||||
if not selected_upstreams:
|
if not selected_upstreams:
|
||||||
return create_error_response(
|
return create_error_response(
|
||||||
@@ -334,6 +491,41 @@ async def _proxy(
|
|||||||
"upstream_error", "All upstreams failed", 502, request=request
|
"upstream_error", "All upstreams failed", 502, request=request
|
||||||
)
|
)
|
||||||
|
|
||||||
|
selector: ModelPathSelector | None = None
|
||||||
|
if MODEL_PATH_HEADER in headers:
|
||||||
|
selector = decode_model_path(headers[MODEL_PATH_HEADER])
|
||||||
|
if (
|
||||||
|
selector is None
|
||||||
|
or sum(
|
||||||
|
name.lower() == MODEL_PATH_HEADER for name, _ in request.headers.items()
|
||||||
|
)
|
||||||
|
!= 1
|
||||||
|
):
|
||||||
|
return create_error_response(
|
||||||
|
"invalid_request",
|
||||||
|
f"Malformed {MODEL_PATH_HEADER} header",
|
||||||
|
400,
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
if not isinstance(model_id, str) or not _model_ids_match(
|
||||||
|
model_id, selector.model_id
|
||||||
|
):
|
||||||
|
return create_error_response(
|
||||||
|
"invalid_request",
|
||||||
|
f"{MODEL_PATH_HEADER} selects model '{selector.model_id}' but the "
|
||||||
|
f"request asks for '{model_id}'",
|
||||||
|
400,
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
if "models" in request_body_dict:
|
||||||
|
return create_error_response(
|
||||||
|
"invalid_request",
|
||||||
|
"Model paths cannot be combined with model fallbacks",
|
||||||
|
400,
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
model_id = selector.model_id
|
||||||
|
|
||||||
candidates = get_candidates(model_id)
|
candidates = get_candidates(model_id)
|
||||||
|
|
||||||
if not candidates:
|
if not candidates:
|
||||||
@@ -341,6 +533,51 @@ async def _proxy(
|
|||||||
"invalid_model", f"Model '{model_id}' not found", 400, request=request
|
"invalid_model", f"Model '{model_id}' not found", 400, request=request
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if selector is not None:
|
||||||
|
pinned = _candidate_for_selector(selector, candidates)
|
||||||
|
if pinned is None:
|
||||||
|
return create_error_response(
|
||||||
|
"invalid_model_path",
|
||||||
|
f"Model '{selector.model_id}' is not routable through provider "
|
||||||
|
f"{selector.provider_id}",
|
||||||
|
404,
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
# Explicit routes must never enter cross-provider failover.
|
||||||
|
candidates = [pinned]
|
||||||
|
|
||||||
|
if selector.endpoint_tag:
|
||||||
|
if (
|
||||||
|
is_ehbp
|
||||||
|
or not request_body_dict
|
||||||
|
or not is_openrouter_base_url(pinned[1].base_url)
|
||||||
|
or _canonical_api_path(path)
|
||||||
|
not in {"chat/completions", "completions", "responses"}
|
||||||
|
):
|
||||||
|
return create_error_response(
|
||||||
|
"unsupported_request",
|
||||||
|
"Endpoint pinning requires an OpenRouter completion or Responses JSON request",
|
||||||
|
400,
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
provider_options = request_body_dict.get("provider", {})
|
||||||
|
if not isinstance(provider_options, dict):
|
||||||
|
return create_error_response(
|
||||||
|
"invalid_request",
|
||||||
|
"provider must be an object",
|
||||||
|
400,
|
||||||
|
request=request,
|
||||||
|
)
|
||||||
|
request_body_dict = {
|
||||||
|
**request_body_dict,
|
||||||
|
"provider": {
|
||||||
|
**provider_options,
|
||||||
|
"order": [selector.endpoint_tag],
|
||||||
|
"allow_fallbacks": False,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
request_body = json.dumps(request_body_dict).encode()
|
||||||
|
|
||||||
if is_ehbp:
|
if is_ehbp:
|
||||||
candidates = [
|
candidates = [
|
||||||
(model, upstream)
|
(model, upstream)
|
||||||
@@ -391,11 +628,21 @@ async def _proxy(
|
|||||||
)
|
)
|
||||||
elif is_responses_api:
|
elif is_responses_api:
|
||||||
return await upstream.handle_x_cashu_responses(
|
return await upstream.handle_x_cashu_responses(
|
||||||
request, x_cashu, path, max_cost_for_model, model_obj
|
request,
|
||||||
|
x_cashu,
|
||||||
|
path,
|
||||||
|
max_cost_for_model,
|
||||||
|
model_obj,
|
||||||
|
request_body=request_body,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return await upstream.handle_x_cashu(
|
return await upstream.handle_x_cashu(
|
||||||
request, x_cashu, path, max_cost_for_model, model_obj
|
request,
|
||||||
|
x_cashu,
|
||||||
|
path,
|
||||||
|
max_cost_for_model,
|
||||||
|
model_obj,
|
||||||
|
request_body=request_body,
|
||||||
)
|
)
|
||||||
except UpstreamError as e:
|
except UpstreamError as e:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -593,17 +840,20 @@ async def _proxy(
|
|||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
# Reactive recovery: some models reject one specific request
|
# Same-provider recovery must not relax an explicit route.
|
||||||
# param (e.g. newer Anthropic models deprecating `temperature`).
|
|
||||||
# When the upstream 400s naming such a param, strip it from the
|
|
||||||
# body and retry the SAME upstream. ``already_stripped`` bounds
|
|
||||||
# this to one retry per distinct param so it always terminates.
|
|
||||||
if response.status_code == 400 and not is_ehbp:
|
if response.status_code == 400 and not is_ehbp:
|
||||||
correction = correct_request(
|
correction = correct_request(
|
||||||
request_body,
|
request_body,
|
||||||
extract_error_message(response),
|
extract_error_message(response),
|
||||||
already_stripped,
|
already_stripped,
|
||||||
)
|
)
|
||||||
|
if correction is not None and selector is not None:
|
||||||
|
corrected_body = json.loads(correction.body)
|
||||||
|
if any(
|
||||||
|
corrected_body.get(field) != request_body_dict.get(field)
|
||||||
|
for field in ("model", "provider")
|
||||||
|
):
|
||||||
|
correction = None
|
||||||
if correction is not None:
|
if correction is not None:
|
||||||
request_body, bad_param = correction.body, correction.label
|
request_body, bad_param = correction.body, correction.label
|
||||||
already_stripped.add(bad_param)
|
already_stripped.add(bad_param)
|
||||||
|
|||||||
@@ -0,0 +1,109 @@
|
|||||||
|
"""In-memory negative cache for terminally failed Cashu token redemptions.
|
||||||
|
|
||||||
|
A dead token (already spent, malformed, zero value) presented as a bearer key
|
||||||
|
triggers a full redemption attempt against the issuing mint on *every* request,
|
||||||
|
because no ``api_keys`` row survives the failed attempt. Polling clients that
|
||||||
|
never back off turn one dead token into thousands of pointless mint calls per
|
||||||
|
day. This cache remembers terminal redemption failures by token hash so
|
||||||
|
repeated presentations are rejected locally with the same error the mint
|
||||||
|
attempt would have produced.
|
||||||
|
|
||||||
|
Only failures whose classification code is in :data:`TERMINAL_REDEMPTION_CODES`
|
||||||
|
are cached — transient failures (mint unreachable, rate-limited, cooldown) must
|
||||||
|
never be cached, or a brief mint outage would poison valid tokens.
|
||||||
|
|
||||||
|
The cache is deliberately in-memory (bounded LRU with TTL) rather than a
|
||||||
|
database row: persisting a row per failed token would let an attacker fill the
|
||||||
|
database with garbage tokens for free.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
|
from collections import OrderedDict
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Callable
|
||||||
|
|
||||||
|
# Redemption ``code`` values that can never succeed on retry. A token that was
|
||||||
|
# already spent, failed to decode, or redeemed to zero value stays that way
|
||||||
|
# forever; swap fees exceeding the token amount only changes if the mint
|
||||||
|
# lowers its fees, which the TTL covers.
|
||||||
|
TERMINAL_REDEMPTION_CODES: frozenset[str] = frozenset(
|
||||||
|
{
|
||||||
|
"cashu_token_already_spent",
|
||||||
|
"invalid_cashu_token",
|
||||||
|
"cashu_token_zero_value",
|
||||||
|
"cashu_token_swap_fees_exceed_amount",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
DEFAULT_MAX_ENTRIES = 10_000
|
||||||
|
DEFAULT_TTL_SECONDS = 24 * 60 * 60
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class CachedRedemptionFailure:
|
||||||
|
"""Sanitized classification of a terminal redemption failure.
|
||||||
|
|
||||||
|
Mirrors the ``(type, status, message, code)`` tuple produced by
|
||||||
|
``classify_redemption_error`` so a cache hit yields a byte-identical error
|
||||||
|
envelope to the original mint-backed failure.
|
||||||
|
"""
|
||||||
|
|
||||||
|
status_code: int
|
||||||
|
error_type: str
|
||||||
|
message: str
|
||||||
|
code: str
|
||||||
|
|
||||||
|
|
||||||
|
class RedemptionNegativeCache:
|
||||||
|
"""Bounded TTL+LRU cache keyed by the SHA-256 hash of the bearer token.
|
||||||
|
|
||||||
|
Not thread-safe by design: all access happens on the asyncio event loop.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
max_entries: int = DEFAULT_MAX_ENTRIES,
|
||||||
|
ttl_seconds: float = DEFAULT_TTL_SECONDS,
|
||||||
|
clock: Callable[[], float] = time.monotonic,
|
||||||
|
) -> None:
|
||||||
|
if max_entries <= 0:
|
||||||
|
raise ValueError("max_entries must be positive")
|
||||||
|
if ttl_seconds <= 0:
|
||||||
|
raise ValueError("ttl_seconds must be positive")
|
||||||
|
self._max_entries = max_entries
|
||||||
|
self._ttl_seconds = ttl_seconds
|
||||||
|
self._clock = clock
|
||||||
|
self._entries: OrderedDict[str, tuple[float, CachedRedemptionFailure]] = (
|
||||||
|
OrderedDict()
|
||||||
|
)
|
||||||
|
|
||||||
|
def get(self, hashed_key: str) -> CachedRedemptionFailure | None:
|
||||||
|
entry = self._entries.get(hashed_key)
|
||||||
|
if entry is None:
|
||||||
|
return None
|
||||||
|
expires_at, failure = entry
|
||||||
|
if self._clock() >= expires_at:
|
||||||
|
del self._entries[hashed_key]
|
||||||
|
return None
|
||||||
|
self._entries.move_to_end(hashed_key)
|
||||||
|
return failure
|
||||||
|
|
||||||
|
def put(self, hashed_key: str, failure: CachedRedemptionFailure) -> None:
|
||||||
|
if hashed_key in self._entries:
|
||||||
|
del self._entries[hashed_key]
|
||||||
|
elif len(self._entries) >= self._max_entries:
|
||||||
|
self._entries.popitem(last=False)
|
||||||
|
self._entries[hashed_key] = (self._clock() + self._ttl_seconds, failure)
|
||||||
|
|
||||||
|
def discard(self, hashed_key: str) -> None:
|
||||||
|
self._entries.pop(hashed_key, None)
|
||||||
|
|
||||||
|
def clear(self) -> None:
|
||||||
|
self._entries.clear()
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self._entries)
|
||||||
|
|
||||||
|
|
||||||
|
# Process-wide singleton used by the bearer-auth path.
|
||||||
|
redemption_negative_cache = RedemptionNegativeCache()
|
||||||
@@ -0,0 +1,613 @@
|
|||||||
|
"""Refund claims: one open payout per API key, recorded before it is paid."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import time
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from fastapi import HTTPException
|
||||||
|
from sqlalchemy.exc import IntegrityError
|
||||||
|
from sqlmodel import col, func, select, update
|
||||||
|
|
||||||
|
from .core.db import (
|
||||||
|
REFUND_OPEN_STATUSES,
|
||||||
|
REFUND_UNRESOLVED_STATUSES,
|
||||||
|
ApiKey,
|
||||||
|
AsyncSession,
|
||||||
|
Refund,
|
||||||
|
create_session,
|
||||||
|
)
|
||||||
|
from .core.db import (
|
||||||
|
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||||
|
)
|
||||||
|
from .core.logging import get_logger
|
||||||
|
from .core.settings import settings
|
||||||
|
from .payment.lnurl import (
|
||||||
|
LNURLError,
|
||||||
|
MeltOutcomeAmbiguousError,
|
||||||
|
MeltUnpaidError,
|
||||||
|
get_lnurl_data,
|
||||||
|
)
|
||||||
|
from .wallet import (
|
||||||
|
check_bolt11_payment_status,
|
||||||
|
is_mint_connection_error,
|
||||||
|
send_to_lnurl,
|
||||||
|
send_token,
|
||||||
|
token_mint_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
RECONCILE_BATCH_LIMIT = 100
|
||||||
|
|
||||||
|
|
||||||
|
def refund_unit(key: ApiKey) -> str:
|
||||||
|
return key.refund_currency or "sat"
|
||||||
|
|
||||||
|
|
||||||
|
def amount_in_unit(amount_msats: int, unit: str) -> int:
|
||||||
|
return amount_msats // 1000 if unit == "sat" else amount_msats
|
||||||
|
|
||||||
|
|
||||||
|
def refund_mint(key: ApiKey) -> str:
|
||||||
|
if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints:
|
||||||
|
return key.refund_mint_url
|
||||||
|
return settings.primary_mint
|
||||||
|
|
||||||
|
|
||||||
|
async def validate_lightning_destination(destination: str) -> None:
|
||||||
|
try:
|
||||||
|
await get_lnurl_data(destination)
|
||||||
|
except LNURLError as e:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400, detail=f"Invalid lightning destination: {e}"
|
||||||
|
)
|
||||||
|
except httpx.HTTPError as e:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400, detail=f"Lightning destination unreachable: {e}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def open_claim(
|
||||||
|
session: AsyncSession,
|
||||||
|
key: ApiKey,
|
||||||
|
*,
|
||||||
|
method: str,
|
||||||
|
destination: str | None,
|
||||||
|
) -> Refund:
|
||||||
|
"""Zero the balance and insert the claim in one transaction."""
|
||||||
|
unit = refund_unit(key)
|
||||||
|
# created_at has second resolution; step past the previous claim so the
|
||||||
|
# newest claim for a key always sorts first.
|
||||||
|
latest = await session.exec(
|
||||||
|
select(func.max(col(Refund.created_at))).where(
|
||||||
|
Refund.api_key_hashed_key == key.hashed_key
|
||||||
|
)
|
||||||
|
)
|
||||||
|
created_at = max(int(time.time()), (latest.one() or 0) + 1)
|
||||||
|
refund = Refund(
|
||||||
|
api_key_hashed_key=key.hashed_key,
|
||||||
|
method=method,
|
||||||
|
destination=destination,
|
||||||
|
amount_msats=key.total_balance,
|
||||||
|
unit=unit,
|
||||||
|
mint_url=refund_mint(key),
|
||||||
|
claimed_at=int(time.time()),
|
||||||
|
created_at=created_at,
|
||||||
|
updated_at=created_at,
|
||||||
|
)
|
||||||
|
debit = (
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
||||||
|
.where(col(ApiKey.balance) == key.balance)
|
||||||
|
.where(col(ApiKey.reserved_balance) == key.reserved_balance)
|
||||||
|
.values(balance=0, reserved_balance=0, reserved_at=None)
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
debited = await session.exec(debit) # type: ignore[call-overload]
|
||||||
|
if debited.rowcount == 0:
|
||||||
|
await session.rollback()
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=409,
|
||||||
|
detail="Balance changed concurrently. Please retry the refund.",
|
||||||
|
)
|
||||||
|
session.add(refund)
|
||||||
|
await session.commit()
|
||||||
|
except IntegrityError:
|
||||||
|
await session.rollback()
|
||||||
|
raise refund_in_progress_error()
|
||||||
|
return refund
|
||||||
|
|
||||||
|
|
||||||
|
async def _close(
|
||||||
|
session: AsyncSession,
|
||||||
|
refund: Refund,
|
||||||
|
*,
|
||||||
|
require_no_quote: bool = False,
|
||||||
|
require_no_token: bool = False,
|
||||||
|
from_statuses: tuple[str, ...] = REFUND_OPEN_STATUSES,
|
||||||
|
**values: object,
|
||||||
|
) -> bool:
|
||||||
|
stmt = (
|
||||||
|
update(Refund)
|
||||||
|
.where(col(Refund.id) == refund.id)
|
||||||
|
.where(col(Refund.status).in_(from_statuses))
|
||||||
|
)
|
||||||
|
if require_no_quote:
|
||||||
|
# A quote recorded since the row was read means a melt may be in flight.
|
||||||
|
stmt = stmt.where(col(Refund.quote_id).is_(None))
|
||||||
|
if require_no_token:
|
||||||
|
stmt = stmt.where(col(Refund.token).is_(None))
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
stmt.values(claimed_at=None, updated_at=int(time.time()), **values)
|
||||||
|
)
|
||||||
|
return bool(result.rowcount)
|
||||||
|
|
||||||
|
|
||||||
|
async def renew_lease(session: AsyncSession, refund: Refund) -> None:
|
||||||
|
"""Push the reconciler lease forward before a slow mint step."""
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
update(Refund)
|
||||||
|
.where(col(Refund.id) == refund.id)
|
||||||
|
.where(col(Refund.status).in_(REFUND_OPEN_STATUSES))
|
||||||
|
.values(claimed_at=int(time.time()))
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def record_quote(refund: Refund, quote_id: str, mint_url: str) -> None:
|
||||||
|
"""Store the quote and its mint before the melt is sent; raises if the claim closed.
|
||||||
|
|
||||||
|
Mint fallback can issue the quote on a different mint than the claim's.
|
||||||
|
"""
|
||||||
|
async with create_session() as session:
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(Refund)
|
||||||
|
.where(col(Refund.id) == refund.id)
|
||||||
|
.where(col(Refund.status).in_(REFUND_OPEN_STATUSES))
|
||||||
|
# Renew the lease so the reconciler leaves the payout alone.
|
||||||
|
.values(
|
||||||
|
quote_id=quote_id,
|
||||||
|
mint_url=mint_url,
|
||||||
|
claimed_at=int(time.time()),
|
||||||
|
updated_at=int(time.time()),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
if not result.rowcount:
|
||||||
|
raise LNURLError("Refund claim closed before the melt was dispatched")
|
||||||
|
refund.quote_id = quote_id
|
||||||
|
refund.mint_url = mint_url
|
||||||
|
|
||||||
|
|
||||||
|
async def settle(
|
||||||
|
session: AsyncSession,
|
||||||
|
refund: Refund,
|
||||||
|
*,
|
||||||
|
quote_id: str | None = None,
|
||||||
|
token: str | None = None,
|
||||||
|
mint_url: str | None = None,
|
||||||
|
) -> bool:
|
||||||
|
values: dict[str, Any] = {"status": "paid"}
|
||||||
|
if quote_id is not None:
|
||||||
|
values["quote_id"] = quote_id
|
||||||
|
if token is not None:
|
||||||
|
values["token"] = token
|
||||||
|
if mint_url is not None:
|
||||||
|
values["mint_url"] = mint_url
|
||||||
|
# The payout side knows the money moved, so a claim the reconciler gave up
|
||||||
|
# on (stuck) is closed as paid too.
|
||||||
|
settled = await _close(
|
||||||
|
session, refund, from_statuses=REFUND_UNRESOLVED_STATUSES, **values
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
if not settled:
|
||||||
|
logger.warning(
|
||||||
|
"refund paid but its claim was already closed",
|
||||||
|
extra={"refund_id": refund.id, "prior_status": refund.status},
|
||||||
|
)
|
||||||
|
return settled
|
||||||
|
|
||||||
|
|
||||||
|
async def release(
|
||||||
|
session: AsyncSession, refund: Refund, *, require_no_quote: bool = False
|
||||||
|
) -> bool:
|
||||||
|
"""Mark the claim failed and restore the balance."""
|
||||||
|
if not await _close(
|
||||||
|
session, refund, require_no_quote=require_no_quote, status="failed"
|
||||||
|
):
|
||||||
|
# Commit, not rollback: rollback after an ORM UPDATE breaks later loads.
|
||||||
|
await session.commit()
|
||||||
|
return False
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == refund.api_key_hashed_key)
|
||||||
|
.values(balance=col(ApiKey.balance) + refund.amount_msats)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
logger.info(
|
||||||
|
"refund released; balance restored",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"key_hash": refund.api_key_hashed_key[:8],
|
||||||
|
"restored_msats": refund.amount_msats,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
async def hold(
|
||||||
|
session: AsyncSession,
|
||||||
|
refund: Refund,
|
||||||
|
quote_id: str | None,
|
||||||
|
*,
|
||||||
|
token: str | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Withhold the balance; the quote or token names what the mint may have paid."""
|
||||||
|
values: dict[str, Any] = {"status": "ambiguous", "quote_id": quote_id}
|
||||||
|
if token is not None:
|
||||||
|
values["token"] = token
|
||||||
|
values["mint_url"] = refund.mint_url
|
||||||
|
await _close(session, refund, **values)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
def refund_in_progress_error(refund: Refund | None = None) -> HTTPException:
|
||||||
|
"""The 409 raised when a key already has an unresolved refund claim."""
|
||||||
|
stuck = refund is not None and refund.status == "stuck"
|
||||||
|
error: dict[str, str] = {
|
||||||
|
"message": (
|
||||||
|
"A refund for this key is unresolved and requires operator reconciliation."
|
||||||
|
if stuck
|
||||||
|
else "A refund for this key is already in progress."
|
||||||
|
),
|
||||||
|
"type": "invalid_request_error",
|
||||||
|
"code": "refund_unresolved" if stuck else "refund_in_progress",
|
||||||
|
}
|
||||||
|
if refund is not None:
|
||||||
|
error["refund_id"] = refund.id
|
||||||
|
error["status"] = refund.status
|
||||||
|
return HTTPException(status_code=409, detail={"error": error})
|
||||||
|
|
||||||
|
|
||||||
|
async def _latest_with_status(
|
||||||
|
session: AsyncSession, key: ApiKey, statuses: tuple[str, ...]
|
||||||
|
) -> Refund | None:
|
||||||
|
result = await session.exec(
|
||||||
|
select(Refund)
|
||||||
|
.where(Refund.api_key_hashed_key == key.hashed_key)
|
||||||
|
.where(col(Refund.status).in_(statuses))
|
||||||
|
.order_by(col(Refund.created_at).desc(), col(Refund.updated_at).desc())
|
||||||
|
)
|
||||||
|
return result.first()
|
||||||
|
|
||||||
|
|
||||||
|
async def latest_open(session: AsyncSession, key: ApiKey) -> Refund | None:
|
||||||
|
"""Latest non-terminal (in-flight) claim for the key, if any.
|
||||||
|
|
||||||
|
An open claim means a prior refund already debited the balance and is still
|
||||||
|
settling, so the balance reads as zero even though a refund is under way.
|
||||||
|
"""
|
||||||
|
return await _latest_with_status(session, key, REFUND_OPEN_STATUSES)
|
||||||
|
|
||||||
|
|
||||||
|
async def latest_stuck(session: AsyncSession, key: ApiKey) -> Refund | None:
|
||||||
|
"""Latest claim the reconciler gave up on; needs operator recovery."""
|
||||||
|
return await _latest_with_status(session, key, ("stuck",))
|
||||||
|
|
||||||
|
|
||||||
|
async def latest_terminal(session: AsyncSession, key: ApiKey) -> Refund | None:
|
||||||
|
"""Latest paid refund of either method.
|
||||||
|
|
||||||
|
Cashu tokens are normally served from cashu_transactions, which tracks
|
||||||
|
collection and sweeping; the claim row is the fallback when that ledger
|
||||||
|
write failed after the token was already issued.
|
||||||
|
"""
|
||||||
|
return await _latest_with_status(session, key, ("paid",))
|
||||||
|
|
||||||
|
|
||||||
|
def describe(refund: Refund) -> dict[str, str]:
|
||||||
|
body: dict[str, str] = {"refund_id": refund.id, "status": refund.status}
|
||||||
|
if refund.token:
|
||||||
|
body["token"] = refund.token
|
||||||
|
if refund.destination:
|
||||||
|
body["recipient"] = refund.destination
|
||||||
|
if refund.unit == "sat":
|
||||||
|
body["sats"] = str(refund.amount_msats // 1000)
|
||||||
|
else:
|
||||||
|
body["msats"] = str(refund.amount_msats)
|
||||||
|
return body
|
||||||
|
|
||||||
|
|
||||||
|
async def _pay_lightning(session: AsyncSession, refund: Refund) -> bool:
|
||||||
|
async def capture_quote(quote: str, mint_url: str) -> None:
|
||||||
|
await record_quote(refund, quote, mint_url)
|
||||||
|
|
||||||
|
try:
|
||||||
|
await send_to_lnurl(
|
||||||
|
amount_in_unit(refund.amount_msats, refund.unit),
|
||||||
|
refund.unit,
|
||||||
|
refund.mint_url,
|
||||||
|
str(refund.destination),
|
||||||
|
on_melt_quote=capture_quote,
|
||||||
|
)
|
||||||
|
except MeltOutcomeAmbiguousError as e:
|
||||||
|
await hold(session, refund, refund.quote_id)
|
||||||
|
logger.error(
|
||||||
|
"refund outcome ambiguous; balance withheld pending reconciliation",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"error": str(e),
|
||||||
|
"key_hash": refund.api_key_hashed_key[:8],
|
||||||
|
"quote_id": refund.quote_id,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
return await settle(session, refund, quote_id=refund.quote_id)
|
||||||
|
|
||||||
|
|
||||||
|
async def _pay_cashu(session: AsyncSession, refund: Refund) -> bool:
|
||||||
|
amount = amount_in_unit(refund.amount_msats, refund.unit)
|
||||||
|
await renew_lease(session, refund)
|
||||||
|
token = await send_token(amount, refund.unit, refund.mint_url)
|
||||||
|
# From here the token is bearer money: keep it on the claim so a failed
|
||||||
|
# settle withholds the balance instead of restoring it.
|
||||||
|
refund.token = token
|
||||||
|
refund.mint_url = token_mint_url(token, refund.mint_url)
|
||||||
|
return await settle(session, refund, token=token, mint_url=refund.mint_url)
|
||||||
|
|
||||||
|
|
||||||
|
async def _record_cashu_payout(refund: Refund) -> None:
|
||||||
|
"""Ledger write for an issued token; the claim row already holds the token,
|
||||||
|
so a failure here must not fail the request or release the balance."""
|
||||||
|
try:
|
||||||
|
await store_cashu_transaction(
|
||||||
|
token=str(refund.token),
|
||||||
|
amount=amount_in_unit(refund.amount_msats, refund.unit),
|
||||||
|
unit=refund.unit,
|
||||||
|
mint_url=refund.mint_url,
|
||||||
|
typ="out",
|
||||||
|
collected=False,
|
||||||
|
source="apikey",
|
||||||
|
api_key_hashed_key=refund.api_key_hashed_key,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
"refund token issued but cashu transaction was not recorded",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"error": str(e),
|
||||||
|
"error_type": type(e).__name__,
|
||||||
|
"key_hash": refund.api_key_hashed_key[:8],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def unresolved_refund_error() -> HTTPException:
|
||||||
|
return HTTPException(
|
||||||
|
status_code=502,
|
||||||
|
detail=(
|
||||||
|
"Refund was dispatched but its outcome is unconfirmed; the "
|
||||||
|
"balance is withheld until reconciliation completes"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _abort(session: AsyncSession, refund: Refund) -> None:
|
||||||
|
"""Fail the claim, or withhold it once a melt quote or token exists.
|
||||||
|
|
||||||
|
A recorded quote means the mint may already have paid; an issued token is
|
||||||
|
already bearer money. In both cases the balance must not be restored.
|
||||||
|
"""
|
||||||
|
if refund.quote_id is None and refund.token is None:
|
||||||
|
await release(session, refund)
|
||||||
|
return
|
||||||
|
await hold(session, refund, refund.quote_id, token=refund.token)
|
||||||
|
logger.error(
|
||||||
|
"refund failed after its payout was dispatched; balance withheld "
|
||||||
|
"pending reconciliation",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"key_hash": refund.api_key_hashed_key[:8],
|
||||||
|
"quote_id": refund.quote_id,
|
||||||
|
"has_token": refund.token is not None,
|
||||||
|
"mint_url": refund.mint_url,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
raise unresolved_refund_error()
|
||||||
|
|
||||||
|
|
||||||
|
async def execute(session: AsyncSession, refund: Refund) -> dict[str, str]:
|
||||||
|
attached_refund = refund
|
||||||
|
# Keep payout evidence outside the identity map: a failed flush/commit can
|
||||||
|
# expire attached attributes, including the only copy of an issued token.
|
||||||
|
refund = Refund(**refund.model_dump())
|
||||||
|
try:
|
||||||
|
if refund.method == "lightning":
|
||||||
|
settled = await _pay_lightning(session, refund)
|
||||||
|
else:
|
||||||
|
settled = await _pay_cashu(session, refund)
|
||||||
|
except MeltOutcomeAmbiguousError:
|
||||||
|
# Already held by _pay_lightning; releasing here would pay out twice.
|
||||||
|
raise unresolved_refund_error()
|
||||||
|
except MeltUnpaidError as e:
|
||||||
|
# The mint answered the melt itself with unpaid: proof that nothing was
|
||||||
|
# sent, so the balance goes back now rather than after reconciliation.
|
||||||
|
await release(session, refund)
|
||||||
|
logger.warning(
|
||||||
|
"refund melt unpaid at the mint; balance restored",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"error": str(e),
|
||||||
|
"key_hash": refund.api_key_hashed_key[:8],
|
||||||
|
"quote_id": refund.quote_id,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=503,
|
||||||
|
detail="Lightning payment failed at the mint; balance restored. Retry later.",
|
||||||
|
)
|
||||||
|
except HTTPException:
|
||||||
|
await session.rollback()
|
||||||
|
await _abort(session, refund)
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
await session.rollback()
|
||||||
|
await _abort(session, refund)
|
||||||
|
logger.error(
|
||||||
|
"refund payout failed",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"error": str(e),
|
||||||
|
"error_type": type(e).__name__,
|
||||||
|
"key_hash": refund.api_key_hashed_key[:8],
|
||||||
|
"method": refund.method,
|
||||||
|
"mint_url": refund.mint_url,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if is_mint_connection_error(e):
|
||||||
|
raise HTTPException(status_code=503, detail="Mint service unavailable")
|
||||||
|
raise HTTPException(status_code=500, detail="Refund failed")
|
||||||
|
|
||||||
|
if refund.method == "cashu":
|
||||||
|
await _record_cashu_payout(refund)
|
||||||
|
|
||||||
|
if settled:
|
||||||
|
refund.status = "paid"
|
||||||
|
refund.claimed_at = None
|
||||||
|
else:
|
||||||
|
# Report the row as it stands rather than a status that was not written.
|
||||||
|
await session.refresh(attached_refund)
|
||||||
|
refund = attached_refund
|
||||||
|
logger.info(
|
||||||
|
"refund paid",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"method": refund.method,
|
||||||
|
"amount_msats": refund.amount_msats,
|
||||||
|
"key_hash": refund.api_key_hashed_key[:8],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return describe(refund)
|
||||||
|
|
||||||
|
|
||||||
|
async def _lease(refund_id: str, now: int, lease_cutoff: int) -> bool:
|
||||||
|
async with create_session() as session:
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(Refund)
|
||||||
|
.where(col(Refund.id) == refund_id)
|
||||||
|
.where(col(Refund.status).in_(REFUND_OPEN_STATUSES))
|
||||||
|
.where(
|
||||||
|
col(Refund.claimed_at).is_(None)
|
||||||
|
| (col(Refund.claimed_at) < lease_cutoff)
|
||||||
|
)
|
||||||
|
.values(claimed_at=now)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return bool(result.rowcount)
|
||||||
|
|
||||||
|
|
||||||
|
async def _reconcile(refund: Refund, now: int) -> None:
|
||||||
|
if refund.method != "lightning":
|
||||||
|
if refund.token is not None:
|
||||||
|
# The token was issued and kept on the claim; the payout is done.
|
||||||
|
async with create_session() as session:
|
||||||
|
await settle(session, refund)
|
||||||
|
return
|
||||||
|
# No quote to query for cashu; withhold the balance and alert once.
|
||||||
|
async with create_session() as session:
|
||||||
|
if await _close(session, refund, require_no_token=True, status="stuck"):
|
||||||
|
await session.commit()
|
||||||
|
logger.critical(
|
||||||
|
"cashu refund stuck; balance withheld, manual reconciliation required",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"key_hash": refund.api_key_hashed_key[:8],
|
||||||
|
"amount_msats": refund.amount_msats,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await session.commit()
|
||||||
|
current = await session.get(Refund, refund.id)
|
||||||
|
if current is not None and current.token is not None:
|
||||||
|
await settle(session, current)
|
||||||
|
return
|
||||||
|
|
||||||
|
if refund.quote_id is None:
|
||||||
|
# Never sent, unless a quote appeared since the row was read.
|
||||||
|
async with create_session() as session:
|
||||||
|
if not await release(session, refund, require_no_quote=True):
|
||||||
|
logger.info(
|
||||||
|
"refund gained a melt quote during reconciliation; left open",
|
||||||
|
extra={"refund_id": refund.id},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
status = await check_bolt11_payment_status(
|
||||||
|
refund.mint_url, refund.unit, refund.quote_id
|
||||||
|
)
|
||||||
|
if status == "paid":
|
||||||
|
async with create_session() as session:
|
||||||
|
await settle(session, refund)
|
||||||
|
elif status == "unpaid":
|
||||||
|
# A fresh melt can report unpaid briefly; trust it only after a timeout.
|
||||||
|
if refund.updated_at > now - settings.refund_claim_timeout_seconds:
|
||||||
|
logger.info(
|
||||||
|
"refund unpaid at the mint but too recent to release; waiting",
|
||||||
|
extra={"refund_id": refund.id, "updated_at": refund.updated_at},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
async with create_session() as session:
|
||||||
|
await release(session, refund)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"refund still unresolved at the mint",
|
||||||
|
extra={"refund_id": refund.id, "melt_status": status},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def reconcile_once() -> None:
|
||||||
|
"""Resolve open claims whose lease has lapsed."""
|
||||||
|
now = int(time.time())
|
||||||
|
lease_cutoff = now - settings.refund_claim_timeout_seconds
|
||||||
|
async with create_session() as session:
|
||||||
|
result = await session.exec(
|
||||||
|
select(Refund)
|
||||||
|
.where(col(Refund.status).in_(REFUND_OPEN_STATUSES))
|
||||||
|
.where(
|
||||||
|
col(Refund.claimed_at).is_(None)
|
||||||
|
| (col(Refund.claimed_at) < lease_cutoff)
|
||||||
|
)
|
||||||
|
.order_by(col(Refund.created_at))
|
||||||
|
.limit(RECONCILE_BATCH_LIMIT)
|
||||||
|
)
|
||||||
|
stale = list(result.all())
|
||||||
|
|
||||||
|
for refund in stale:
|
||||||
|
if not await _lease(refund.id, now, lease_cutoff):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
await _reconcile(refund, now)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
"refund reconciliation failed",
|
||||||
|
extra={
|
||||||
|
"refund_id": refund.id,
|
||||||
|
"error": str(e),
|
||||||
|
"error_type": type(e).__name__,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def periodic_refund_reconcile() -> None:
|
||||||
|
while True:
|
||||||
|
await asyncio.sleep(settings.refund_reconcile_interval_seconds)
|
||||||
|
try:
|
||||||
|
await reconcile_once()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
"refund reconcile loop error",
|
||||||
|
extra={"error": str(e), "error_type": type(e).__name__},
|
||||||
|
)
|
||||||
+678
-53
@@ -15,9 +15,6 @@ from ..core.db import (
|
|||||||
UpstreamProviderRow,
|
UpstreamProviderRow,
|
||||||
create_session,
|
create_session,
|
||||||
)
|
)
|
||||||
from ..core.db import (
|
|
||||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
|
||||||
)
|
|
||||||
from ..payment.price import sats_usd_price
|
from ..payment.price import sats_usd_price
|
||||||
from ..wallet import (
|
from ..wallet import (
|
||||||
Bolt11PaymentAmbiguous,
|
Bolt11PaymentAmbiguous,
|
||||||
@@ -27,7 +24,7 @@ from ..wallet import (
|
|||||||
maximum_owner_cashu_balance_sats,
|
maximum_owner_cashu_balance_sats,
|
||||||
prepare_bolt11_payment,
|
prepare_bolt11_payment,
|
||||||
release_token_reservation,
|
release_token_reservation,
|
||||||
send_token,
|
send_token_from_owner_locked,
|
||||||
token_mint_url,
|
token_mint_url,
|
||||||
wallet_operation_guard,
|
wallet_operation_guard,
|
||||||
)
|
)
|
||||||
@@ -49,6 +46,7 @@ PPQ_PHASES = frozenset({PPQ_PHASE_CLAIMED, PPQ_PHASE_IN_FLIGHT, PPQ_PHASE_RECONC
|
|||||||
PPQ_SETTLEMENT_ATTEMPTS = 5
|
PPQ_SETTLEMENT_ATTEMPTS = 5
|
||||||
PPQ_SETTLEMENT_POLL_SECONDS = 2
|
PPQ_SETTLEMENT_POLL_SECONDS = 2
|
||||||
PPQ_PENDING_TTL_SECONDS = 15 * 60
|
PPQ_PENDING_TTL_SECONDS = 15 * 60
|
||||||
|
PPQ_SETTLED_COOLDOWN_SECONDS = 5 * 60
|
||||||
PPQ_MAX_INVOICE_PREMIUM = 1.10
|
PPQ_MAX_INVOICE_PREMIUM = 1.10
|
||||||
PPQ_MIN_TOPUP_USD = 1
|
PPQ_MIN_TOPUP_USD = 1
|
||||||
PPQ_MAX_TOPUP_USD = 500
|
PPQ_MAX_TOPUP_USD = 500
|
||||||
@@ -58,6 +56,33 @@ PPQ_MAX_TOPUP_USD = 500
|
|||||||
# the damage instead of letting the worker drain the owner's mint funds one
|
# the damage instead of letting the worker drain the owner's mint funds one
|
||||||
# per-transaction-capped payment at a time.
|
# per-transaction-capped payment at a time.
|
||||||
PPQ_MAX_DAILY_TOPUP_USD = 300
|
PPQ_MAX_DAILY_TOPUP_USD = 300
|
||||||
|
# Routstr-to-Routstr claim lifecycle. "claimed" holds the slot while the token
|
||||||
|
# is being minted; nothing has left the wallet yet. "sent" means a bearer token
|
||||||
|
# was handed to the peer and only the peer's balance can say whether it landed.
|
||||||
|
# "backoff" holds the failure count between attempts, and "halted" stops the
|
||||||
|
# provider entirely until an admin releases it.
|
||||||
|
ROUTSTR_PHASE_CLAIMED = "claimed"
|
||||||
|
ROUTSTR_PHASE_SENT = "sent"
|
||||||
|
ROUTSTR_PHASE_BACKOFF = "backoff"
|
||||||
|
ROUTSTR_PHASE_HALTED = "halted"
|
||||||
|
ROUTSTR_PHASES = frozenset(
|
||||||
|
{
|
||||||
|
ROUTSTR_PHASE_CLAIMED,
|
||||||
|
ROUTSTR_PHASE_SENT,
|
||||||
|
ROUTSTR_PHASE_BACKOFF,
|
||||||
|
ROUTSTR_PHASE_HALTED,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
ROUTSTR_PENDING_TTL_SECONDS = 15 * 60
|
||||||
|
ROUTSTR_BACKOFF_BASE_SECONDS = 15 * 60
|
||||||
|
ROUTSTR_MAX_TOPUP_FAILURES = 3
|
||||||
|
ROUTSTR_MIN_TOPUP_SATS = 1
|
||||||
|
ROUTSTR_MAX_TOPUP_SATS = 1_000_000
|
||||||
|
# Rolling 24h ceiling on total Routstr auto top-up spend across all peers. The
|
||||||
|
# per-attempt claim already stops a peer from being paid twice for the same
|
||||||
|
# uncredited token; this bounds the total even when every attempt is credited
|
||||||
|
# and the peer simply keeps reporting a below-threshold balance.
|
||||||
|
ROUTSTR_MAX_DAILY_TOPUP_SATS = 2_000_000
|
||||||
|
|
||||||
|
|
||||||
async def periodic_auto_topup() -> None:
|
async def periodic_auto_topup() -> None:
|
||||||
@@ -72,7 +97,7 @@ async def periodic_auto_topup() -> None:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(
|
logger.error(
|
||||||
"Auto top-up cycle failed",
|
"Auto top-up cycle failed",
|
||||||
extra={"error": str(e), "error_type": type(e).__name__},
|
extra={"error": repr(e), "error_type": type(e).__name__},
|
||||||
)
|
)
|
||||||
|
|
||||||
await asyncio.sleep(AUTO_TOPUP_INTERVAL_SECONDS)
|
await asyncio.sleep(AUTO_TOPUP_INTERVAL_SECONDS)
|
||||||
@@ -103,7 +128,7 @@ async def _run_auto_topup_cycle() -> None:
|
|||||||
extra={
|
extra={
|
||||||
"provider_id": row.id,
|
"provider_id": row.id,
|
||||||
"base_url": row.base_url,
|
"base_url": row.base_url,
|
||||||
"error": str(e),
|
"error": repr(e),
|
||||||
"error_type": type(e).__name__,
|
"error_type": type(e).__name__,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
@@ -134,12 +159,16 @@ async def _reconcile_all_ppq_claims() -> set[int]:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(
|
logger.error(
|
||||||
"PPQ claim reconciliation failed",
|
"PPQ claim reconciliation failed",
|
||||||
extra={"provider_id": row.id, "error": str(e)},
|
extra={
|
||||||
|
"provider_id": row.id,
|
||||||
|
"error": repr(e),
|
||||||
|
"error_type": type(e).__name__,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
return active_provider_ids
|
return active_provider_ids
|
||||||
|
|
||||||
|
|
||||||
def _invalid_ppq_number(value: object, *, integer: bool = False) -> bool:
|
def _invalid_topup_number(value: object, *, integer: bool = False) -> bool:
|
||||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||||
return True
|
return True
|
||||||
try:
|
try:
|
||||||
@@ -160,9 +189,9 @@ def validate_ppq_auto_topup_settings(settings: dict | None) -> str | None:
|
|||||||
|
|
||||||
threshold = settings.get("topup_threshold")
|
threshold = settings.get("topup_threshold")
|
||||||
amount = settings.get("topup_amount_limit")
|
amount = settings.get("topup_amount_limit")
|
||||||
if _invalid_ppq_number(threshold):
|
if _invalid_topup_number(threshold):
|
||||||
return "PPQ auto top-up threshold must be a positive number"
|
return "PPQ auto top-up threshold must be a positive number"
|
||||||
if _invalid_ppq_number(amount, integer=True):
|
if _invalid_topup_number(amount, integer=True):
|
||||||
return "PPQ auto top-up amount must be a positive whole number"
|
return "PPQ auto top-up amount must be a positive whole number"
|
||||||
amount_usd = int(typing.cast(int | float, amount))
|
amount_usd = int(typing.cast(int | float, amount))
|
||||||
if not PPQ_MIN_TOPUP_USD <= amount_usd <= PPQ_MAX_TOPUP_USD:
|
if not PPQ_MIN_TOPUP_USD <= amount_usd <= PPQ_MAX_TOPUP_USD:
|
||||||
@@ -173,6 +202,62 @@ def validate_ppq_auto_topup_settings(settings: dict | None) -> str | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
_legacy_threshold_hinted: set[int] = set()
|
||||||
|
|
||||||
|
|
||||||
|
def _routstr_threshold_sats(row: UpstreamProviderRow, settings: dict) -> float:
|
||||||
|
"""Balance, in sats, below which a top-up fires.
|
||||||
|
|
||||||
|
``topup_threshold_sats`` is used as written. A legacy ``topup_threshold``
|
||||||
|
keeps the thousandfold it has always been compared with — reinterpreting it
|
||||||
|
as sats would drop an operator's trigger point by a factor of 1000 on
|
||||||
|
upgrade and leave the peer to run dry. The hint names the value to migrate
|
||||||
|
to, once per provider rather than once per scheduler tick.
|
||||||
|
"""
|
||||||
|
explicit = settings.get("topup_threshold_sats")
|
||||||
|
if explicit is not None:
|
||||||
|
return float(typing.cast(int | float, explicit))
|
||||||
|
|
||||||
|
threshold_sats = float(typing.cast(int | float, settings["topup_threshold"])) * 1000
|
||||||
|
if row.id is not None and row.id not in _legacy_threshold_hinted:
|
||||||
|
_legacy_threshold_hinted.add(row.id)
|
||||||
|
logger.warning(
|
||||||
|
"Routstr auto top-up uses the legacy unitless threshold; set "
|
||||||
|
"topup_threshold_sats to state the unit",
|
||||||
|
extra={"provider_id": row.id, "threshold_sats": threshold_sats},
|
||||||
|
)
|
||||||
|
return threshold_sats
|
||||||
|
|
||||||
|
|
||||||
|
def validate_routstr_auto_topup_settings(settings: dict | None) -> str | None:
|
||||||
|
"""Return why enabled Routstr auto top-up settings are invalid, if anything."""
|
||||||
|
if not settings or not settings.get("auto_topup"):
|
||||||
|
return None
|
||||||
|
|
||||||
|
threshold_sats = settings.get("topup_threshold_sats")
|
||||||
|
legacy_threshold = settings.get("topup_threshold")
|
||||||
|
amount = settings.get("topup_amount_limit")
|
||||||
|
mint_url = settings.get("topup_mint_url")
|
||||||
|
if threshold_sats is not None:
|
||||||
|
if _invalid_topup_number(threshold_sats):
|
||||||
|
return "Routstr auto top-up threshold must be a positive number of sats"
|
||||||
|
elif legacy_threshold is None:
|
||||||
|
return "Routstr auto top-up requires a threshold"
|
||||||
|
elif _invalid_topup_number(legacy_threshold):
|
||||||
|
return "Routstr auto top-up threshold must be a positive number"
|
||||||
|
if _invalid_topup_number(amount, integer=True):
|
||||||
|
return "Routstr auto top-up amount must be a positive whole number"
|
||||||
|
amount_sats = int(typing.cast(int | float, amount))
|
||||||
|
if not ROUTSTR_MIN_TOPUP_SATS <= amount_sats <= ROUTSTR_MAX_TOPUP_SATS:
|
||||||
|
return (
|
||||||
|
f"Routstr auto top-up amount must be between {ROUTSTR_MIN_TOPUP_SATS} "
|
||||||
|
f"and {ROUTSTR_MAX_TOPUP_SATS} sats"
|
||||||
|
)
|
||||||
|
if not isinstance(mint_url, str) or not mint_url.strip():
|
||||||
|
return "Routstr auto top-up requires a mint URL"
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
async def _check_and_topup_ppq_from_row(row: UpstreamProviderRow) -> None:
|
async def _check_and_topup_ppq_from_row(row: UpstreamProviderRow) -> None:
|
||||||
settings: dict = {}
|
settings: dict = {}
|
||||||
if row.provider_settings:
|
if row.provider_settings:
|
||||||
@@ -212,22 +297,18 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
|||||||
if not settings.get("auto_topup"):
|
if not settings.get("auto_topup"):
|
||||||
return
|
return
|
||||||
|
|
||||||
threshold = settings.get("topup_threshold")
|
problem = validate_routstr_auto_topup_settings(settings)
|
||||||
amount = settings.get("topup_amount_limit")
|
if problem is not None:
|
||||||
mint_url = settings.get("topup_mint_url")
|
|
||||||
|
|
||||||
if not threshold or not amount or not mint_url:
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Auto top-up enabled but missing configuration",
|
"Auto top-up enabled but its configuration is invalid",
|
||||||
extra={
|
extra={"provider_id": row.id, "problem": problem},
|
||||||
"provider_id": row.id,
|
|
||||||
"has_threshold": bool(threshold),
|
|
||||||
"has_amount": bool(amount),
|
|
||||||
"has_mint": bool(mint_url),
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
threshold_sats = _routstr_threshold_sats(row, settings)
|
||||||
|
amount = int(settings["topup_amount_limit"])
|
||||||
|
mint_url = str(settings["topup_mint_url"])
|
||||||
|
|
||||||
if not row.api_key:
|
if not row.api_key:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -235,16 +316,40 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
|||||||
provider = RoutstrUpstreamProvider.from_db_row(row)
|
provider = RoutstrUpstreamProvider.from_db_row(row)
|
||||||
if provider is None:
|
if provider is None:
|
||||||
return
|
return
|
||||||
|
if await _reconcile_routstr_state(row, provider):
|
||||||
|
return
|
||||||
|
|
||||||
balance = await provider.get_balance()
|
balance = await provider.get_balance()
|
||||||
|
|
||||||
if balance is None:
|
if balance is None or not math.isfinite(balance) or balance < 0:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Could not fetch balance for auto top-up",
|
"Could not fetch balance for auto top-up",
|
||||||
extra={"provider_id": row.id, "base_url": row.base_url},
|
extra={"provider_id": row.id, "base_url": row.base_url},
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
if balance >= threshold * 1000:
|
if balance >= threshold_sats:
|
||||||
|
return
|
||||||
|
|
||||||
|
spent_24h_sats = await _routstr_spent_last_24h_sats()
|
||||||
|
if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS:
|
||||||
|
logger.critical(
|
||||||
|
"Auto top-up skipped: rolling 24h spend cap reached",
|
||||||
|
extra={
|
||||||
|
"provider_id": row.id,
|
||||||
|
"spent_24h_sats": spent_24h_sats,
|
||||||
|
"topup_amount": amount,
|
||||||
|
"daily_cap_sats": ROUTSTR_MAX_DAILY_TOPUP_SATS,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# The balance the peer must report before another token may be sent. Any
|
||||||
|
# shortfall is treated as "not credited": the token is a bearer instrument
|
||||||
|
# and a peer that took one without crediting it must not be handed another.
|
||||||
|
expected_sats = math.floor(balance) + amount
|
||||||
|
operation_id = await _claim_routstr_topup(row, expected_sats=expected_sats)
|
||||||
|
if operation_id is None:
|
||||||
return
|
return
|
||||||
|
|
||||||
# Balance is below threshold - create token and top up
|
# Balance is below threshold - create token and top up
|
||||||
@@ -253,40 +358,33 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
|||||||
extra={
|
extra={
|
||||||
"provider_id": row.id,
|
"provider_id": row.id,
|
||||||
"balance": balance,
|
"balance": balance,
|
||||||
"threshold": threshold,
|
"threshold_sats": threshold_sats,
|
||||||
"topup_amount": amount,
|
"topup_amount": amount,
|
||||||
"mint_url": mint_url,
|
"mint_url": mint_url,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
token = await send_token(amount, "sat", mint_url)
|
async with wallet_operation_guard():
|
||||||
except Exception as e:
|
# Keep the spend cap and audit mutation in one wallet lock.
|
||||||
logger.error(
|
spent_24h_sats = await _routstr_spent_last_24h_sats()
|
||||||
"Failed to create cashu token for auto top-up",
|
if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS:
|
||||||
extra={
|
raise ValueError("Routstr auto top-up daily spend cap reached")
|
||||||
"provider_id": row.id,
|
token = await send_token_from_owner_locked(amount, "sat", mint_url)
|
||||||
"amount": amount,
|
|
||||||
"mint_url": mint_url,
|
|
||||||
"error": str(e),
|
|
||||||
},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
actual_mint_url = token_mint_url(token, mint_url)
|
actual_mint_url = token_mint_url(token, mint_url)
|
||||||
try:
|
try:
|
||||||
await store_cashu_transaction(
|
await _persist_routstr_token_and_mark_sent(
|
||||||
|
row,
|
||||||
|
operation_id,
|
||||||
|
expected_sats=expected_sats,
|
||||||
token=token,
|
token=token,
|
||||||
amount=amount,
|
amount=amount,
|
||||||
unit="sat",
|
|
||||||
mint_url=actual_mint_url,
|
mint_url=actual_mint_url,
|
||||||
typ="out",
|
|
||||||
collected=False,
|
|
||||||
source="auto_topup",
|
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.critical(
|
logger.critical(
|
||||||
"Aborting auto top-up because its cashu token could not be persisted",
|
"Aborting auto top-up because its token and sent claim "
|
||||||
|
"could not be persisted atomically",
|
||||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
@@ -297,7 +395,7 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
|||||||
extra={
|
extra={
|
||||||
"provider_id": row.id,
|
"provider_id": row.id,
|
||||||
"mint_url": actual_mint_url,
|
"mint_url": actual_mint_url,
|
||||||
"error": str(error),
|
"error": repr(error),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -305,6 +403,19 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
|||||||
"Auto-topup token was released after persistence failed",
|
"Auto-topup token was released after persistence failed",
|
||||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||||
)
|
)
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to create or persist cashu token for auto top-up",
|
||||||
|
extra={
|
||||||
|
"provider_id": row.id,
|
||||||
|
"amount": amount,
|
||||||
|
"mint_url": mint_url,
|
||||||
|
"error": repr(e),
|
||||||
|
"error_type": type(e).__name__,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
await _release_routstr_claim(row, operation_id)
|
||||||
return
|
return
|
||||||
|
|
||||||
result = await provider.topup(token)
|
result = await provider.topup(token)
|
||||||
@@ -348,6 +459,481 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _routstr_state_id(row: UpstreamProviderRow) -> str:
|
||||||
|
if row.id is None:
|
||||||
|
raise ValueError("Routstr auto top-up requires a persisted provider row")
|
||||||
|
return _routstr_state_id_for_provider(row.id)
|
||||||
|
|
||||||
|
|
||||||
|
def _routstr_state_id_for_provider(provider_id: int | str) -> str:
|
||||||
|
return f"routstr-auto-topup-{provider_id}"
|
||||||
|
|
||||||
|
|
||||||
|
class RoutstrClaim(typing.NamedTuple):
|
||||||
|
operation_id: str
|
||||||
|
# Worker lease for "claimed"/"sent", retry-not-before for "backoff", and
|
||||||
|
# meaningless for "halted".
|
||||||
|
deadline: int
|
||||||
|
phase: str
|
||||||
|
# Peer balance in sats that proves this attempt was credited.
|
||||||
|
expected_sats: int
|
||||||
|
failures: int
|
||||||
|
|
||||||
|
|
||||||
|
def _routstr_request_id(
|
||||||
|
operation_id: str,
|
||||||
|
deadline: int,
|
||||||
|
phase: str,
|
||||||
|
expected_sats: int,
|
||||||
|
failures: int,
|
||||||
|
) -> str:
|
||||||
|
return f"routstr:{operation_id}:{deadline}:{phase}:{expected_sats}:{failures}"
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_routstr_request_id(request_id: str | None) -> RoutstrClaim | None:
|
||||||
|
parts = (request_id or "").split(":", 5)
|
||||||
|
if len(parts) != 6 or parts[0] != "routstr" or parts[3] not in ROUTSTR_PHASES:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
deadline = int(parts[2])
|
||||||
|
expected_sats = int(parts[4])
|
||||||
|
failures = int(parts[5])
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return None
|
||||||
|
return RoutstrClaim(parts[1], deadline, parts[3], expected_sats, failures)
|
||||||
|
|
||||||
|
|
||||||
|
async def _routstr_spent_last_24h_sats() -> int:
|
||||||
|
"""Total sats committed to Routstr auto top-ups in the last 24 hours.
|
||||||
|
|
||||||
|
Uncollected rows count too: a token whose delivery is unconfirmed is spent
|
||||||
|
for capping purposes. Rows marked ``collected=False, swept=True`` record
|
||||||
|
tokens that were provably returned to the wallet and are excluded.
|
||||||
|
"""
|
||||||
|
cutoff = int(time.time()) - 24 * 60 * 60
|
||||||
|
async with create_session() as session:
|
||||||
|
rows = (
|
||||||
|
await session.exec(
|
||||||
|
select(CashuTransaction.amount, CashuTransaction.unit).where(
|
||||||
|
col(CashuTransaction.source) == "auto_topup",
|
||||||
|
col(CashuTransaction.type) == "out",
|
||||||
|
col(CashuTransaction.created_at) >= cutoff,
|
||||||
|
or_(
|
||||||
|
col(CashuTransaction.collected) == True, # noqa: E712
|
||||||
|
col(CashuTransaction.swept) == False, # noqa: E712
|
||||||
|
),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
return sum(
|
||||||
|
amount if unit == "sat" else math.ceil(amount / 1000) for amount, unit in rows
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _routstr_provider_is_claimable(
|
||||||
|
session: AsyncSession, provider_id: int | str | None
|
||||||
|
) -> bool:
|
||||||
|
"""Re-read the provider inside the claim transaction.
|
||||||
|
|
||||||
|
Same reasoning as :func:`_ppq_provider_is_claimable`: a provider deleted or
|
||||||
|
retyped concurrently must either be visible here or lose the race against
|
||||||
|
the claim we are about to write.
|
||||||
|
"""
|
||||||
|
if provider_id is None:
|
||||||
|
return False
|
||||||
|
current = await session.get(UpstreamProviderRow, provider_id)
|
||||||
|
return current is not None and current.provider_type == "routstr"
|
||||||
|
|
||||||
|
|
||||||
|
async def _claim_routstr_topup(
|
||||||
|
row: UpstreamProviderRow, *, expected_sats: int
|
||||||
|
) -> str | None:
|
||||||
|
"""Acquire the provider's single durable auto top-up slot.
|
||||||
|
|
||||||
|
An expired backoff hands its failure count to the new attempt, so repeated
|
||||||
|
non-crediting peers still walk towards the halt instead of resetting the
|
||||||
|
counter every cycle.
|
||||||
|
"""
|
||||||
|
state_id = _routstr_state_id(row)
|
||||||
|
operation_id = uuid.uuid4().hex
|
||||||
|
deadline = int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
if not await _routstr_provider_is_claimable(session, row.id):
|
||||||
|
return None
|
||||||
|
existing = await session.get(CashuTransaction, state_id)
|
||||||
|
if existing is not None:
|
||||||
|
failures = 0
|
||||||
|
if not (existing.collected or existing.swept):
|
||||||
|
claim = _parse_routstr_request_id(existing.request_id)
|
||||||
|
if (
|
||||||
|
claim is None
|
||||||
|
or claim.phase != ROUTSTR_PHASE_BACKOFF
|
||||||
|
or time.time() < claim.deadline
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
failures = claim.failures
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(CashuTransaction)
|
||||||
|
.where(
|
||||||
|
col(CashuTransaction.id) == state_id,
|
||||||
|
# Fence on the exact row that was read: any concurrent
|
||||||
|
# writer that moved the claim must win instead of us.
|
||||||
|
col(CashuTransaction.request_id) == existing.request_id,
|
||||||
|
)
|
||||||
|
.values(
|
||||||
|
token="pending",
|
||||||
|
amount=0,
|
||||||
|
unit="sat",
|
||||||
|
mint_url=None,
|
||||||
|
request_id=_routstr_request_id(
|
||||||
|
operation_id,
|
||||||
|
deadline,
|
||||||
|
ROUTSTR_PHASE_CLAIMED,
|
||||||
|
expected_sats,
|
||||||
|
failures,
|
||||||
|
),
|
||||||
|
collected=False,
|
||||||
|
swept=False,
|
||||||
|
created_at=int(time.time()),
|
||||||
|
source="routstr_auto_topup_claim",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
if (getattr(result, "rowcount", 0) or 0) != 1:
|
||||||
|
return None
|
||||||
|
return operation_id
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with create_session() as session:
|
||||||
|
if not await _routstr_provider_is_claimable(session, row.id):
|
||||||
|
return None
|
||||||
|
session.add(
|
||||||
|
CashuTransaction(
|
||||||
|
id=state_id,
|
||||||
|
token="pending",
|
||||||
|
amount=0,
|
||||||
|
unit="sat",
|
||||||
|
type="out",
|
||||||
|
request_id=_routstr_request_id(
|
||||||
|
operation_id,
|
||||||
|
deadline,
|
||||||
|
ROUTSTR_PHASE_CLAIMED,
|
||||||
|
expected_sats,
|
||||||
|
0,
|
||||||
|
),
|
||||||
|
collected=False,
|
||||||
|
source="routstr_auto_topup_claim",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
except IntegrityError:
|
||||||
|
return None
|
||||||
|
return operation_id
|
||||||
|
|
||||||
|
|
||||||
|
async def _advance_routstr_claim(
|
||||||
|
row: UpstreamProviderRow,
|
||||||
|
operation_id: str,
|
||||||
|
*,
|
||||||
|
deadline: int,
|
||||||
|
phase: str,
|
||||||
|
expected_sats: int,
|
||||||
|
failures: int,
|
||||||
|
token: str | None = None,
|
||||||
|
amount: int | None = None,
|
||||||
|
mint_url: str | None = None,
|
||||||
|
) -> bool:
|
||||||
|
"""Move this worker's claim to another phase, if it still owns it."""
|
||||||
|
values: dict[str, object] = {
|
||||||
|
"request_id": _routstr_request_id(
|
||||||
|
operation_id, deadline, phase, expected_sats, failures
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if token is not None:
|
||||||
|
values.update(token=token, amount=amount, mint_url=mint_url)
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(CashuTransaction)
|
||||||
|
.where(
|
||||||
|
col(CashuTransaction.id) == _routstr_state_id(row),
|
||||||
|
col(CashuTransaction.request_id).like(f"routstr:{operation_id}:%"),
|
||||||
|
col(CashuTransaction.collected) == False, # noqa: E712
|
||||||
|
col(CashuTransaction.swept) == False, # noqa: E712
|
||||||
|
)
|
||||||
|
.values(**values)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return (getattr(result, "rowcount", 0) or 0) == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def _set_routstr_state_terminal(
|
||||||
|
row: UpstreamProviderRow, operation_id: str, *, collected: bool, swept: bool
|
||||||
|
) -> bool:
|
||||||
|
"""Finish an attempt only if this worker still owns the claim."""
|
||||||
|
async with create_session() as session:
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(CashuTransaction)
|
||||||
|
.where(
|
||||||
|
col(CashuTransaction.id) == _routstr_state_id(row),
|
||||||
|
col(CashuTransaction.request_id).like(f"routstr:{operation_id}:%"),
|
||||||
|
col(CashuTransaction.collected) == False, # noqa: E712
|
||||||
|
col(CashuTransaction.swept) == False, # noqa: E712
|
||||||
|
)
|
||||||
|
.values(collected=collected, swept=swept)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return (getattr(result, "rowcount", 0) or 0) == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def _release_routstr_claim(row: UpstreamProviderRow, operation_id: str) -> None:
|
||||||
|
"""Hand back a claim whose token never left the wallet."""
|
||||||
|
if not await _set_routstr_state_terminal(
|
||||||
|
row, operation_id, collected=False, swept=True
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"Could not release the auto top-up claim after a pre-send failure; "
|
||||||
|
"it is owned by another attempt",
|
||||||
|
extra={"provider_id": row.id},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _persist_routstr_token_and_mark_sent(
|
||||||
|
row: UpstreamProviderRow,
|
||||||
|
operation_id: str,
|
||||||
|
*,
|
||||||
|
expected_sats: int,
|
||||||
|
token: str,
|
||||||
|
amount: int,
|
||||||
|
mint_url: str,
|
||||||
|
) -> None:
|
||||||
|
state_id = _routstr_state_id(row)
|
||||||
|
async with create_session() as session:
|
||||||
|
state = await session.get(CashuTransaction, state_id)
|
||||||
|
claim = _parse_routstr_request_id(state.request_id if state else None)
|
||||||
|
if (
|
||||||
|
state is None
|
||||||
|
or state.collected
|
||||||
|
or state.swept
|
||||||
|
or claim is None
|
||||||
|
or claim.operation_id != operation_id
|
||||||
|
or claim.phase != ROUTSTR_PHASE_CLAIMED
|
||||||
|
):
|
||||||
|
raise RuntimeError("Routstr auto top-up claim ownership was lost")
|
||||||
|
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(CashuTransaction)
|
||||||
|
.where(
|
||||||
|
col(CashuTransaction.id) == state_id,
|
||||||
|
col(CashuTransaction.request_id) == state.request_id,
|
||||||
|
col(CashuTransaction.collected) == False, # noqa: E712
|
||||||
|
col(CashuTransaction.swept) == False, # noqa: E712
|
||||||
|
)
|
||||||
|
.values(
|
||||||
|
request_id=_routstr_request_id(
|
||||||
|
operation_id,
|
||||||
|
int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS,
|
||||||
|
ROUTSTR_PHASE_SENT,
|
||||||
|
expected_sats,
|
||||||
|
claim.failures,
|
||||||
|
),
|
||||||
|
token=token,
|
||||||
|
amount=amount,
|
||||||
|
unit="sat",
|
||||||
|
mint_url=mint_url,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if (getattr(result, "rowcount", 0) or 0) != 1:
|
||||||
|
await session.rollback()
|
||||||
|
raise RuntimeError("Routstr auto top-up claim ownership was lost")
|
||||||
|
|
||||||
|
session.add(
|
||||||
|
CashuTransaction(
|
||||||
|
id=uuid.uuid4().hex,
|
||||||
|
token=token,
|
||||||
|
amount=amount,
|
||||||
|
unit="sat",
|
||||||
|
mint_url=mint_url,
|
||||||
|
type="out",
|
||||||
|
collected=False,
|
||||||
|
source="auto_topup",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def _current_routstr_claim(row: UpstreamProviderRow) -> RoutstrClaim | None:
|
||||||
|
async with create_session() as session:
|
||||||
|
transaction = await session.get(CashuTransaction, _routstr_state_id(row))
|
||||||
|
if transaction is None:
|
||||||
|
return None
|
||||||
|
return _parse_routstr_request_id(transaction.request_id)
|
||||||
|
|
||||||
|
|
||||||
|
async def _reconcile_routstr_state(
|
||||||
|
row: UpstreamProviderRow, provider: RoutstrUpstreamProvider
|
||||||
|
) -> bool:
|
||||||
|
"""Return True while a prior attempt must suppress a new payment."""
|
||||||
|
async with create_session() as session:
|
||||||
|
transaction = await session.get(CashuTransaction, _routstr_state_id(row))
|
||||||
|
if transaction is None or transaction.collected or transaction.swept:
|
||||||
|
return False
|
||||||
|
|
||||||
|
claim = _parse_routstr_request_id(transaction.request_id)
|
||||||
|
if claim is None:
|
||||||
|
logger.critical(
|
||||||
|
"Malformed auto top-up state; suppressing duplicate payment",
|
||||||
|
extra={"provider_id": row.id},
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
now = time.time()
|
||||||
|
if claim.phase == ROUTSTR_PHASE_HALTED:
|
||||||
|
return True
|
||||||
|
if claim.phase == ROUTSTR_PHASE_BACKOFF:
|
||||||
|
return now < claim.deadline
|
||||||
|
if claim.phase == ROUTSTR_PHASE_CLAIMED:
|
||||||
|
# Nothing left the wallet, so a dead worker's slot is free to reuse.
|
||||||
|
if now < claim.deadline:
|
||||||
|
return True
|
||||||
|
return not await _set_routstr_state_terminal(
|
||||||
|
row, claim.operation_id, collected=False, swept=True
|
||||||
|
)
|
||||||
|
|
||||||
|
balance = await provider.get_balance()
|
||||||
|
if (
|
||||||
|
balance is not None
|
||||||
|
and math.isfinite(balance)
|
||||||
|
and balance >= claim.expected_sats
|
||||||
|
):
|
||||||
|
if not await _set_routstr_state_terminal(
|
||||||
|
row, claim.operation_id, collected=True, swept=False
|
||||||
|
):
|
||||||
|
logger.critical(
|
||||||
|
"Auto top-up was credited but its claim was already released; "
|
||||||
|
"a duplicate top-up is possible on the next cycle",
|
||||||
|
extra={"provider_id": row.id},
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
if now < claim.deadline:
|
||||||
|
return True
|
||||||
|
|
||||||
|
failures = claim.failures + 1
|
||||||
|
if failures >= ROUTSTR_MAX_TOPUP_FAILURES:
|
||||||
|
await _advance_routstr_claim(
|
||||||
|
row,
|
||||||
|
claim.operation_id,
|
||||||
|
deadline=claim.deadline,
|
||||||
|
phase=ROUTSTR_PHASE_HALTED,
|
||||||
|
expected_sats=claim.expected_sats,
|
||||||
|
failures=failures,
|
||||||
|
)
|
||||||
|
logger.critical(
|
||||||
|
"Auto top-up halted: the peer repeatedly failed to credit a token",
|
||||||
|
extra={
|
||||||
|
"provider_id": row.id,
|
||||||
|
"base_url": row.base_url,
|
||||||
|
"failures": failures,
|
||||||
|
"admin_action": (
|
||||||
|
f"POST /admin/api/upstream-providers/{row.id}"
|
||||||
|
"/routstr-auto-topup/release"
|
||||||
|
),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
await _advance_routstr_claim(
|
||||||
|
row,
|
||||||
|
claim.operation_id,
|
||||||
|
deadline=int(now) + ROUTSTR_BACKOFF_BASE_SECONDS * 2 ** (failures - 1),
|
||||||
|
phase=ROUTSTR_PHASE_BACKOFF,
|
||||||
|
expected_sats=claim.expected_sats,
|
||||||
|
failures=failures,
|
||||||
|
)
|
||||||
|
logger.warning(
|
||||||
|
"Auto top-up was not credited by the peer; backing off",
|
||||||
|
extra={
|
||||||
|
"provider_id": row.id,
|
||||||
|
"base_url": row.base_url,
|
||||||
|
"expected_sats": claim.expected_sats,
|
||||||
|
"failures": failures,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
async def get_routstr_auto_topup_state(provider_id: int) -> dict[str, object]:
|
||||||
|
"""Return admin-safe state for a provider's durable Routstr claim."""
|
||||||
|
async with create_session() as session:
|
||||||
|
transaction = await session.get(
|
||||||
|
CashuTransaction, _routstr_state_id_for_provider(provider_id)
|
||||||
|
)
|
||||||
|
if transaction is None or transaction.collected or transaction.swept:
|
||||||
|
return {"active": False}
|
||||||
|
|
||||||
|
claim = _parse_routstr_request_id(transaction.request_id)
|
||||||
|
return {
|
||||||
|
"active": True,
|
||||||
|
# Echoed back verbatim on release so a claim that moved on since the
|
||||||
|
# admin reviewed it fails the write instead of being swept unseen.
|
||||||
|
"state_token": transaction.request_id,
|
||||||
|
"operation_id": claim.operation_id if claim else None,
|
||||||
|
"phase": claim.phase if claim else None,
|
||||||
|
"expected_sats": claim.expected_sats if claim else None,
|
||||||
|
"failures": claim.failures if claim else None,
|
||||||
|
"deadline": claim.deadline if claim else None,
|
||||||
|
"created_at": transaction.created_at,
|
||||||
|
"amount": transaction.amount,
|
||||||
|
"unit": transaction.unit,
|
||||||
|
"mint_url": transaction.mint_url,
|
||||||
|
"malformed": claim is None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class RoutstrReleaseOutcome(typing.NamedTuple):
|
||||||
|
released: bool
|
||||||
|
reason: str
|
||||||
|
|
||||||
|
|
||||||
|
async def release_routstr_auto_topup_state(
|
||||||
|
provider_id: int, *, state_token: str | None
|
||||||
|
) -> RoutstrReleaseOutcome:
|
||||||
|
"""Clear a halted or stuck claim after an admin reconciles the peer.
|
||||||
|
|
||||||
|
Unlike a Lightning melt there is no in-flight window to protect: the token
|
||||||
|
is already with the peer or still in the wallet either way. What the fence
|
||||||
|
does protect is the admin's decision — the row must be byte-identical to
|
||||||
|
the one they reviewed, so a claim that advanced in the meantime is not
|
||||||
|
swept on the strength of stale information.
|
||||||
|
"""
|
||||||
|
state_id = _routstr_state_id_for_provider(provider_id)
|
||||||
|
async with create_session() as session:
|
||||||
|
transaction = await session.get(CashuTransaction, state_id)
|
||||||
|
|
||||||
|
if transaction is None or transaction.collected or transaction.swept:
|
||||||
|
return RoutstrReleaseOutcome(False, "no_active_claim")
|
||||||
|
if transaction.request_id != state_token:
|
||||||
|
return RoutstrReleaseOutcome(False, "stale_state")
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
|
update(CashuTransaction)
|
||||||
|
.where(
|
||||||
|
col(CashuTransaction.id) == state_id,
|
||||||
|
col(CashuTransaction.request_id) == state_token,
|
||||||
|
col(CashuTransaction.collected) == False, # noqa: E712
|
||||||
|
col(CashuTransaction.swept) == False, # noqa: E712
|
||||||
|
)
|
||||||
|
.values(swept=True)
|
||||||
|
)
|
||||||
|
if (getattr(result, "rowcount", 0) or 0) == 1:
|
||||||
|
await session.commit()
|
||||||
|
return RoutstrReleaseOutcome(True, "released")
|
||||||
|
await session.rollback()
|
||||||
|
return RoutstrReleaseOutcome(False, "claim_changed")
|
||||||
|
|
||||||
|
|
||||||
def _ppq_state_id(row: UpstreamProviderRow) -> str:
|
def _ppq_state_id(row: UpstreamProviderRow) -> str:
|
||||||
if row.id is None:
|
if row.id is None:
|
||||||
raise ValueError("PPQ auto top-up requires a persisted provider row")
|
raise ValueError("PPQ auto top-up requires a persisted provider row")
|
||||||
@@ -528,7 +1114,11 @@ async def _set_ppq_state_terminal(
|
|||||||
col(CashuTransaction.collected) == False, # noqa: E712
|
col(CashuTransaction.collected) == False, # noqa: E712
|
||||||
col(CashuTransaction.swept) == False, # noqa: E712
|
col(CashuTransaction.swept) == False, # noqa: E712
|
||||||
)
|
)
|
||||||
.values(collected=collected, swept=swept)
|
.values(
|
||||||
|
collected=collected,
|
||||||
|
swept=swept,
|
||||||
|
created_at=int(time.time()) if collected else CashuTransaction.created_at,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
updated = (getattr(result, "rowcount", 0) or 0) == 1
|
updated = (getattr(result, "rowcount", 0) or 0) == 1
|
||||||
if updated:
|
if updated:
|
||||||
@@ -553,8 +1143,10 @@ async def _reconcile_ppq_state(
|
|||||||
"""
|
"""
|
||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
transaction = await session.get(CashuTransaction, _ppq_state_id(row))
|
transaction = await session.get(CashuTransaction, _ppq_state_id(row))
|
||||||
if transaction is None or transaction.collected or transaction.swept:
|
if transaction is None or transaction.swept:
|
||||||
return False
|
return False
|
||||||
|
if transaction.collected:
|
||||||
|
return int(time.time()) - transaction.created_at < PPQ_SETTLED_COOLDOWN_SECONDS
|
||||||
|
|
||||||
claim = _parse_ppq_request_id(transaction.request_id)
|
claim = _parse_ppq_request_id(transaction.request_id)
|
||||||
if claim is None:
|
if claim is None:
|
||||||
@@ -620,7 +1212,7 @@ async def _reconcile_ppq_state(
|
|||||||
|
|
||||||
|
|
||||||
async def _ppq_provider_is_claimable(
|
async def _ppq_provider_is_claimable(
|
||||||
session: AsyncSession, provider_id: int | None
|
session: AsyncSession, row: UpstreamProviderRow
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""Re-read the provider inside the claim transaction.
|
"""Re-read the provider inside the claim transaction.
|
||||||
|
|
||||||
@@ -631,10 +1223,16 @@ async def _ppq_provider_is_claimable(
|
|||||||
this the worker could create a claim for a provider that no longer
|
this the worker could create a claim for a provider that no longer
|
||||||
exists, orphaning it forever.
|
exists, orphaning it forever.
|
||||||
"""
|
"""
|
||||||
if provider_id is None:
|
if row.id is None:
|
||||||
return False
|
return False
|
||||||
current = await session.get(UpstreamProviderRow, provider_id)
|
current = await session.get(UpstreamProviderRow, row.id)
|
||||||
return current is not None and current.provider_type == "ppqai"
|
return bool(
|
||||||
|
current is not None
|
||||||
|
and current.enabled
|
||||||
|
and current.provider_type == "ppqai"
|
||||||
|
and current.api_key == row.api_key
|
||||||
|
and current.provider_settings == row.provider_settings
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
||||||
@@ -645,10 +1243,17 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
|||||||
request_id = _ppq_request_id(operation_id, expires_at, PPQ_PHASE_CLAIMED, "pending")
|
request_id = _ppq_request_id(operation_id, expires_at, PPQ_PHASE_CLAIMED, "pending")
|
||||||
|
|
||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
if not await _ppq_provider_is_claimable(session, row.id):
|
if not await _ppq_provider_is_claimable(session, row):
|
||||||
return None
|
return None
|
||||||
existing = await session.get(CashuTransaction, state_id)
|
existing = await session.get(CashuTransaction, state_id)
|
||||||
if existing is not None:
|
if existing is not None:
|
||||||
|
if (
|
||||||
|
existing.collected
|
||||||
|
and not existing.swept
|
||||||
|
and int(time.time()) - existing.created_at
|
||||||
|
< PPQ_SETTLED_COOLDOWN_SECONDS
|
||||||
|
):
|
||||||
|
return None
|
||||||
result = await session.exec( # type: ignore[call-overload]
|
result = await session.exec( # type: ignore[call-overload]
|
||||||
update(CashuTransaction)
|
update(CashuTransaction)
|
||||||
.where(
|
.where(
|
||||||
@@ -679,7 +1284,7 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
|||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
# Same fencing as the update path: the provider must still exist
|
# Same fencing as the update path: the provider must still exist
|
||||||
# inside the transaction that creates the claim.
|
# inside the transaction that creates the claim.
|
||||||
if not await _ppq_provider_is_claimable(session, row.id):
|
if not await _ppq_provider_is_claimable(session, row):
|
||||||
return None
|
return None
|
||||||
session.add(
|
session.add(
|
||||||
CashuTransaction(
|
CashuTransaction(
|
||||||
@@ -900,6 +1505,26 @@ async def _check_and_topup_ppq(row: UpstreamProviderRow, settings: dict) -> None
|
|||||||
if balance >= threshold_usd:
|
if balance >= threshold_usd:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# Require two low-balance reads before creating an invoice.
|
||||||
|
confirmed_balance = await provider.get_balance()
|
||||||
|
if (
|
||||||
|
confirmed_balance is None
|
||||||
|
or not math.isfinite(confirmed_balance)
|
||||||
|
or confirmed_balance < 0
|
||||||
|
or confirmed_balance >= threshold_usd
|
||||||
|
):
|
||||||
|
logger.info(
|
||||||
|
"PPQ auto top-up aborted by balance confirmation",
|
||||||
|
extra={
|
||||||
|
"provider_id": row.id,
|
||||||
|
"first_balance_usd": balance,
|
||||||
|
"confirmed_balance_usd": confirmed_balance,
|
||||||
|
"threshold_usd": threshold_usd,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
balance = confirmed_balance
|
||||||
|
|
||||||
# Perform local pricing and owner-funds checks before asking PPQ to create
|
# Perform local pricing and owner-funds checks before asking PPQ to create
|
||||||
# an invoice. The exact mint quote still has to be checked afterward, but
|
# an invoice. The exact mint quote still has to be checked afterward, but
|
||||||
# predictable local failures should not leave abandoned PPQ invoices.
|
# predictable local failures should not leave abandoned PPQ invoices.
|
||||||
|
|||||||
+575
-279
File diff suppressed because it is too large
Load Diff
@@ -23,7 +23,11 @@ import litellm
|
|||||||
from fastapi.responses import Response
|
from fastapi.responses import Response
|
||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..payment.helpers import estimate_tokens
|
from ..payment.helpers import (
|
||||||
|
_count_prompt_token_ids,
|
||||||
|
estimate_prompt_tokens,
|
||||||
|
estimate_tokens,
|
||||||
|
)
|
||||||
from ..payment.models import Model
|
from ..payment.models import Model
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
@@ -39,11 +43,66 @@ def _parse_request_body(request_body: bytes | None) -> dict[str, Any]:
|
|||||||
return parsed if isinstance(parsed, dict) else {}
|
return parsed if isinstance(parsed, dict) else {}
|
||||||
|
|
||||||
|
|
||||||
def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
|
def _model_name(model_obj: Model | None, body: dict[str, Any]) -> str:
|
||||||
|
if model_obj is not None:
|
||||||
|
return model_obj.forwarded_model_id or model_obj.id or ""
|
||||||
|
body_model = body.get("model")
|
||||||
|
return body_model if isinstance(body_model, str) else ""
|
||||||
|
|
||||||
|
|
||||||
|
def _count_with_litellm(
|
||||||
|
model: str, body: dict[str, Any], include_legacy_prompt: bool = False
|
||||||
|
) -> int:
|
||||||
messages = body.get("messages")
|
messages = body.get("messages")
|
||||||
if not isinstance(messages, list):
|
if not isinstance(messages, list):
|
||||||
messages = []
|
messages = []
|
||||||
|
|
||||||
|
if "input" in body:
|
||||||
|
response_input = body["input"]
|
||||||
|
if isinstance(response_input, str):
|
||||||
|
messages = [{"role": "user", "content": response_input}]
|
||||||
|
elif isinstance(response_input, list):
|
||||||
|
messages = []
|
||||||
|
for item in response_input:
|
||||||
|
if not isinstance(item, dict) or "role" not in item:
|
||||||
|
raise ValueError(
|
||||||
|
"Responses input requires fallback token estimation"
|
||||||
|
)
|
||||||
|
content = item.get("content", "")
|
||||||
|
if isinstance(content, list):
|
||||||
|
parts = []
|
||||||
|
for part in content:
|
||||||
|
if not isinstance(part, dict) or part.get("type") not in (
|
||||||
|
"input_text",
|
||||||
|
"output_text",
|
||||||
|
"text",
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"Non-text Responses input requires fallback token estimation"
|
||||||
|
)
|
||||||
|
parts.append({"type": "text", "text": part.get("text", "")})
|
||||||
|
content = parts
|
||||||
|
messages.append({"role": item["role"], "content": content})
|
||||||
|
else:
|
||||||
|
raise ValueError("Unsupported Responses input")
|
||||||
|
if body.get("instructions"):
|
||||||
|
messages.insert(0, {"role": "system", "content": body["instructions"]})
|
||||||
|
|
||||||
|
prompt_token_ids = 0
|
||||||
|
if include_legacy_prompt:
|
||||||
|
prompt = body.get("prompt")
|
||||||
|
if isinstance(prompt, str):
|
||||||
|
prompt_texts = [prompt]
|
||||||
|
elif isinstance(prompt, list):
|
||||||
|
prompt_texts = [item for item in prompt if isinstance(item, str)]
|
||||||
|
else:
|
||||||
|
prompt_texts = []
|
||||||
|
prompt_token_ids = _count_prompt_token_ids(prompt)
|
||||||
|
messages = [
|
||||||
|
*({"role": "user", "content": text} for text in prompt_texts if text),
|
||||||
|
*messages,
|
||||||
|
]
|
||||||
|
|
||||||
system = body.get("system")
|
system = body.get("system")
|
||||||
if isinstance(system, str) and system:
|
if isinstance(system, str) and system:
|
||||||
messages = [{"role": "system", "content": system}, *messages]
|
messages = [{"role": "system", "content": system}, *messages]
|
||||||
@@ -58,7 +117,7 @@ def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
|
|||||||
|
|
||||||
tools = body.get("tools") if isinstance(body.get("tools"), list) else None
|
tools = body.get("tools") if isinstance(body.get("tools"), list) else None
|
||||||
|
|
||||||
return int(
|
return prompt_token_ids + int(
|
||||||
litellm.token_counter(
|
litellm.token_counter(
|
||||||
model=model,
|
model=model,
|
||||||
messages=messages,
|
messages=messages,
|
||||||
@@ -67,6 +126,173 @@ def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _count_text_with_litellm(model: str, text: str) -> int:
|
||||||
|
return int(
|
||||||
|
litellm.token_counter(
|
||||||
|
model=model,
|
||||||
|
text=text,
|
||||||
|
count_response_tokens=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _generated_text(value: object) -> list[str]:
|
||||||
|
"""Extract generated text/tool arguments without counting response metadata."""
|
||||||
|
generated_keys = {
|
||||||
|
"arguments",
|
||||||
|
"content",
|
||||||
|
"delta",
|
||||||
|
"output_text",
|
||||||
|
"partial_json",
|
||||||
|
"reasoning",
|
||||||
|
"reasoning_content",
|
||||||
|
"text",
|
||||||
|
"thinking",
|
||||||
|
}
|
||||||
|
parts: list[str] = []
|
||||||
|
|
||||||
|
def walk(item: object, key: str | None = None) -> None:
|
||||||
|
if isinstance(item, str):
|
||||||
|
if key in generated_keys:
|
||||||
|
parts.append(item)
|
||||||
|
return
|
||||||
|
if isinstance(item, list):
|
||||||
|
for child in item:
|
||||||
|
walk(child, key)
|
||||||
|
return
|
||||||
|
if isinstance(item, dict):
|
||||||
|
for child_key, child in item.items():
|
||||||
|
walk(child, child_key)
|
||||||
|
|
||||||
|
walk(value)
|
||||||
|
return parts
|
||||||
|
|
||||||
|
|
||||||
|
class MissingUsageEstimator:
|
||||||
|
"""Estimate billable usage when an upstream omits its usage trailer.
|
||||||
|
|
||||||
|
The reservation is deliberately absent from this class: it is an
|
||||||
|
authorization ceiling, not an input to usage measurement.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, request_body: bytes | None, model_obj: Model | None) -> None:
|
||||||
|
self.body = _parse_request_body(request_body)
|
||||||
|
self.model_name = _model_name(model_obj, self.body)
|
||||||
|
self._output_parts: list[str] = []
|
||||||
|
self._input_tokens: int | None = None
|
||||||
|
|
||||||
|
def _estimate_input_tokens(self) -> int:
|
||||||
|
if self._input_tokens is not None:
|
||||||
|
return self._input_tokens
|
||||||
|
try:
|
||||||
|
self._input_tokens = _count_with_litellm(
|
||||||
|
self.model_name, self.body, include_legacy_prompt=True
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
self._input_tokens = estimate_prompt_tokens(self.body)
|
||||||
|
logger.debug(
|
||||||
|
"litellm request token count failed; using local estimator",
|
||||||
|
extra={
|
||||||
|
"model": self.model_name,
|
||||||
|
"error": str(exc),
|
||||||
|
"error_type": type(exc).__name__,
|
||||||
|
"estimated_tokens": self._input_tokens,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return self._input_tokens
|
||||||
|
|
||||||
|
@property
|
||||||
|
def output_text(self) -> str:
|
||||||
|
return "".join(self._output_parts)
|
||||||
|
|
||||||
|
def observe(self, response_data: object) -> None:
|
||||||
|
if isinstance(response_data, dict):
|
||||||
|
event_type = response_data.get("type")
|
||||||
|
if event_type in ("response.completed", "response.incomplete"):
|
||||||
|
response = response_data.get("response")
|
||||||
|
if isinstance(response, dict) and isinstance(
|
||||||
|
response.get("output"), list
|
||||||
|
):
|
||||||
|
# Terminal output is a snapshot, not another text delta.
|
||||||
|
self._output_parts = _generated_text(response["output"])
|
||||||
|
return
|
||||||
|
if isinstance(event_type, str) and event_type.endswith(".done"):
|
||||||
|
# Responses API ``*.done`` events repeat text already streamed
|
||||||
|
# via ``*.delta`` events; counting both would double-bill.
|
||||||
|
return
|
||||||
|
self._output_parts.extend(_generated_text(response_data))
|
||||||
|
|
||||||
|
def estimated_usage(self, model: str | None = None) -> dict[str, Any] | None:
|
||||||
|
"""Local usage estimate, or None when the upstream generated no text."""
|
||||||
|
if not self.output_text:
|
||||||
|
return None
|
||||||
|
return self.response_data(model)["usage"]
|
||||||
|
|
||||||
|
def billing_data(
|
||||||
|
self,
|
||||||
|
response_data: dict[str, Any] | None,
|
||||||
|
model: str | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Use measured usage when present, otherwise return a local estimate."""
|
||||||
|
if isinstance(response_data, dict):
|
||||||
|
usage = response_data.get("usage")
|
||||||
|
if not isinstance(usage, dict):
|
||||||
|
nested = response_data.get("response")
|
||||||
|
usage = nested.get("usage") if isinstance(nested, dict) else None
|
||||||
|
if isinstance(usage, dict) and usage:
|
||||||
|
return {
|
||||||
|
"model": model or response_data.get("model") or self.model_name,
|
||||||
|
"usage": usage,
|
||||||
|
}
|
||||||
|
if not self._output_parts:
|
||||||
|
self.observe(response_data)
|
||||||
|
return self.response_data(model)
|
||||||
|
|
||||||
|
def response_data(self, model: str | None = None) -> dict[str, Any]:
|
||||||
|
text = self.output_text
|
||||||
|
try:
|
||||||
|
output_tokens = (
|
||||||
|
_count_text_with_litellm(self.model_name, text) if text else 0
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
output_tokens = len(text) // 3
|
||||||
|
logger.debug(
|
||||||
|
"litellm response token count failed; using local estimator",
|
||||||
|
extra={
|
||||||
|
"model": self.model_name,
|
||||||
|
"error": str(exc),
|
||||||
|
"error_type": type(exc).__name__,
|
||||||
|
"estimated_tokens": output_tokens,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
input_tokens = max(0, int(self._estimate_input_tokens()))
|
||||||
|
output_tokens = max(0, int(output_tokens))
|
||||||
|
return {
|
||||||
|
"model": model or self.model_name or "unknown",
|
||||||
|
"usage": {
|
||||||
|
"input_tokens": input_tokens,
|
||||||
|
"output_tokens": output_tokens,
|
||||||
|
"total_tokens": input_tokens + output_tokens,
|
||||||
|
"estimated": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
def openai_response_data(self, model: str | None = None) -> dict[str, Any]:
|
||||||
|
"""Same estimate in the OpenAI chat-completions usage dialect."""
|
||||||
|
data = self.response_data(model)
|
||||||
|
usage = data["usage"]
|
||||||
|
return {
|
||||||
|
"model": data["model"],
|
||||||
|
"usage": {
|
||||||
|
"prompt_tokens": usage["input_tokens"],
|
||||||
|
"completion_tokens": usage["output_tokens"],
|
||||||
|
"total_tokens": usage["total_tokens"],
|
||||||
|
"estimated": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def count_tokens_locally(
|
def count_tokens_locally(
|
||||||
request_body: bytes | None,
|
request_body: bytes | None,
|
||||||
model_obj: Model | None,
|
model_obj: Model | None,
|
||||||
@@ -75,13 +301,7 @@ def count_tokens_locally(
|
|||||||
touching the upstream. Always returns 200; never raises."""
|
touching the upstream. Always returns 200; never raises."""
|
||||||
body = _parse_request_body(request_body)
|
body = _parse_request_body(request_body)
|
||||||
|
|
||||||
model_name = ""
|
model_name = _model_name(model_obj, body)
|
||||||
if model_obj is not None:
|
|
||||||
model_name = model_obj.forwarded_model_id or model_obj.id or ""
|
|
||||||
if not model_name:
|
|
||||||
body_model = body.get("model")
|
|
||||||
if isinstance(body_model, str):
|
|
||||||
model_name = body_model
|
|
||||||
|
|
||||||
input_tokens: int
|
input_tokens: int
|
||||||
try:
|
try:
|
||||||
|
|||||||
+365
-230
@@ -10,17 +10,17 @@ from urllib.parse import urlsplit, urlunsplit
|
|||||||
|
|
||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
from fastapi.responses import Response, StreamingResponse
|
from fastapi.responses import Response, StreamingResponse
|
||||||
from sqlalchemy import case
|
|
||||||
from sqlmodel import col, update
|
|
||||||
|
|
||||||
from ..auth import (
|
from ..auth import (
|
||||||
ROUTSTR_FEE_PERCENT,
|
ROUTSTR_FEE_PERCENT,
|
||||||
ReservationSnapshot,
|
ReservationSnapshot,
|
||||||
|
_charge_reservation_rows,
|
||||||
_claim_reservation_for_charge,
|
_claim_reservation_for_charge,
|
||||||
|
_stop_reservation_heartbeat,
|
||||||
_validate_reservation_snapshot,
|
_validate_reservation_snapshot,
|
||||||
get_billing_key,
|
|
||||||
get_reservation_snapshot,
|
get_reservation_snapshot,
|
||||||
payments_logger,
|
payments_logger,
|
||||||
|
release_reservation,
|
||||||
)
|
)
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.db import (
|
from ..core.db import (
|
||||||
@@ -31,7 +31,7 @@ from ..core.db import (
|
|||||||
from ..core.db import (
|
from ..core.db import (
|
||||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||||
)
|
)
|
||||||
from ..core.exceptions import UpstreamError
|
from ..core.exceptions import EhbpTimeoutError, UpstreamError
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
from ..payment.cost_calculation import (
|
from ..payment.cost_calculation import (
|
||||||
CostData,
|
CostData,
|
||||||
@@ -40,8 +40,13 @@ from ..payment.cost_calculation import (
|
|||||||
)
|
)
|
||||||
from ..payment.helpers import create_error_response
|
from ..payment.helpers import create_error_response
|
||||||
from ..payment.models import Model
|
from ..payment.models import Model
|
||||||
from ..wallet import recieve_token, send_token
|
from ..wallet import (
|
||||||
from .tinfoil_trailer import forward_with_trailer
|
SPENT_TOKEN_CODES,
|
||||||
|
classify_redemption_error,
|
||||||
|
recieve_token,
|
||||||
|
send_token,
|
||||||
|
)
|
||||||
|
from .tinfoil_trailer import TrailerResponse, forward_with_trailer
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -57,6 +62,55 @@ _TINFOIL_ALLOWED_ENCLAVE_HOST_SUFFIX = ".tinfoil.sh"
|
|||||||
_TINFOIL_ALLOWED_ENCLAVE_HOSTS = frozenset({"tinfoil.sh"})
|
_TINFOIL_ALLOWED_ENCLAVE_HOSTS = frozenset({"tinfoil.sh"})
|
||||||
|
|
||||||
|
|
||||||
|
_KEY_CONFIG_PROBLEM_TYPE = "urn:ietf:params:ehbp:error:key-config"
|
||||||
|
|
||||||
|
|
||||||
|
def _is_ehbp_key_config_response(resp: TrailerResponse) -> bool:
|
||||||
|
"""Check whether an upstream EHBP response is a key-config mismatch.
|
||||||
|
|
||||||
|
The enclave returns ``422 application/problem+json`` with
|
||||||
|
``type=urn:ietf:params:ehbp:error:key-config`` when it cannot decrypt the
|
||||||
|
request body — meaning the client's HPKE key is stale (the enclave rotated
|
||||||
|
keys). The proxy must pass this response through with its original
|
||||||
|
content type so EHBP clients can detect it and trigger re-attestation.
|
||||||
|
"""
|
||||||
|
if resp.status_code != 422:
|
||||||
|
return False
|
||||||
|
ct = ""
|
||||||
|
for k, v in resp.headers:
|
||||||
|
if k.lower() == "content-type":
|
||||||
|
ct = v.lower()
|
||||||
|
break
|
||||||
|
media_type = ct.split(";", 1)[0].strip()
|
||||||
|
if media_type != "application/problem+json":
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
body = json.loads(resp.body)
|
||||||
|
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||||
|
return False
|
||||||
|
return isinstance(body, dict) and body.get("type") == _KEY_CONFIG_PROBLEM_TYPE
|
||||||
|
|
||||||
|
|
||||||
|
def _passthrough_key_config_response(resp: TrailerResponse) -> Response:
|
||||||
|
"""Return the enclave's key-config 422 with its original body and content
|
||||||
|
type so the EHBP client's ``KeyConfigMismatchError`` detection fires.
|
||||||
|
|
||||||
|
Only the content type is forwarded. ``Ehbp-Response-Nonce`` must be
|
||||||
|
dropped: a nonce only carries meaning for an *encrypted* response body,
|
||||||
|
and the stock ``ehbp`` client (``shouldDecryptResponse``) checks for the
|
||||||
|
nonce *before* checking for a key-config mismatch — forwarding it would
|
||||||
|
send that client down the decrypt path on this plaintext error body, so
|
||||||
|
the re-attestation loop would never fire. Content-length is recomputed
|
||||||
|
from the body, and upstream-internal headers are filtered out.
|
||||||
|
"""
|
||||||
|
return Response(
|
||||||
|
content=resp.body,
|
||||||
|
status_code=422,
|
||||||
|
headers={"content-type": "application/problem+json"},
|
||||||
|
media_type="application/problem+json",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _normalize_upstream_model_id(model_id: str | None) -> str:
|
def _normalize_upstream_model_id(model_id: str | None) -> str:
|
||||||
"""Normalize casing and whitespace for upstream identity comparisons."""
|
"""Normalize casing and whitespace for upstream identity comparisons."""
|
||||||
if not model_id:
|
if not model_id:
|
||||||
@@ -73,28 +127,43 @@ _PROXY_ONLY_HEADERS = frozenset(
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Namespace prefix the routstr catalog applies to Tinfoil models
|
||||||
|
# (e.g. ``tinfoil-deepseek-v4-1-flash``). The SDK strips this prefix for the
|
||||||
|
# encrypted body (``getTinfoilUpstreamModelId`` in client/TinfoilSecure.ts), so
|
||||||
|
# the enclave always reports the *bare* upstream model id in the usage-metrics
|
||||||
|
# header even though the routstr model id and ``forwarded_model_id`` carry it.
|
||||||
|
TINFOIL_MODEL_PREFIX = "tinfoil-"
|
||||||
|
|
||||||
|
|
||||||
def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None:
|
def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None:
|
||||||
"""Parse ``X-Tinfoil-Usage-Metrics`` into an OpenAI-style usage dict.
|
"""Parse ``X-Tinfoil-Usage-Metrics`` into an OpenAI-style usage dict.
|
||||||
|
|
||||||
The header format is::
|
The header format is::
|
||||||
|
|
||||||
prompt=<n>,completion=<n>,total=<n>[,model=<name>]
|
prompt=<n>,completion=<n>,total=<n>[,cached_prompt_tokens=<n>,
|
||||||
|
uncached_prompt_tokens=<n>][,model=<name>][,cost_usd=<usd>]
|
||||||
|
|
||||||
|
``prompt`` is the inclusive prompt total and ``cached_prompt_tokens`` is
|
||||||
|
the cache-read portion included within it. Routstr maps these to
|
||||||
|
``prompt_tokens`` and ``cache_read_input_tokens`` so ``normalize_usage``
|
||||||
|
can subtract the cached read from the prompt total (OpenAI-family
|
||||||
|
semantics). ``cost_usd`` is parsed as a float and kept for logging/
|
||||||
|
cross-checking only — billing uses the token path.
|
||||||
|
|
||||||
The ``model`` field (added in tinfoilsh/confidential-model-router PR #385)
|
The ``model`` field (added in tinfoilsh/confidential-model-router PR #385)
|
||||||
is extracted as a string and included in the returned dict under the
|
is extracted as a string so callers can compare the served model against
|
||||||
``"model"`` key so callers can compare the served model against the
|
the requested one and adjust pricing.
|
||||||
requested one and adjust pricing.
|
|
||||||
|
|
||||||
Returns a dict like ``{"prompt_tokens": n, "completion_tokens": n,
|
Returns a dict suitable for :func:`calculate_cost`, or ``None`` when the
|
||||||
"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.
|
header is absent or malformed.
|
||||||
"""
|
"""
|
||||||
if not header_value:
|
if not header_value:
|
||||||
return None
|
return None
|
||||||
parts: dict[str, int] = {}
|
|
||||||
|
int_parts: dict[str, int] = {}
|
||||||
model: str | None = None
|
model: str | None = None
|
||||||
|
cost_usd: float | None = None
|
||||||
|
|
||||||
for item in header_value.split(","):
|
for item in header_value.split(","):
|
||||||
key, sep, value = item.partition("=")
|
key, sep, value = item.partition("=")
|
||||||
if not sep:
|
if not sep:
|
||||||
@@ -104,31 +173,45 @@ def parse_tinfoil_usage_metrics(header_value: str | None) -> dict | None:
|
|||||||
if key == "model":
|
if key == "model":
|
||||||
model = value
|
model = value
|
||||||
continue
|
continue
|
||||||
|
if key == "cost_usd":
|
||||||
try:
|
try:
|
||||||
parts[key] = int(value)
|
cost_usd = float(value)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
cost_usd = None
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
int_parts[key] = int(value)
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
continue
|
continue
|
||||||
prompt = parts.get("prompt")
|
|
||||||
completion = parts.get("completion")
|
prompt = int_parts.get("prompt")
|
||||||
if prompt is not None and completion is not None:
|
completion = int_parts.get("completion")
|
||||||
result: dict[str, int | str] = {
|
if prompt is None or completion is None:
|
||||||
"prompt_tokens": prompt,
|
|
||||||
"completion_tokens": completion,
|
|
||||||
}
|
|
||||||
if "total" in parts:
|
|
||||||
result["total_tokens"] = parts["total"]
|
|
||||||
if model:
|
|
||||||
result["model"] = model
|
|
||||||
return result
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Failed to parse X-Tinfoil-Usage-Metrics header",
|
"Failed to parse X-Tinfoil-Usage-Metrics header",
|
||||||
extra={
|
extra={
|
||||||
"header_value": header_value,
|
"header_value": header_value,
|
||||||
"parsed_parts": parts,
|
"parsed_parts": int_parts,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
result: dict[str, int | float | str] = {
|
||||||
|
"prompt_tokens": prompt,
|
||||||
|
"completion_tokens": completion,
|
||||||
|
}
|
||||||
|
if "total" in int_parts:
|
||||||
|
result["total_tokens"] = int_parts["total"]
|
||||||
|
if "cached_prompt_tokens" in int_parts:
|
||||||
|
result["cache_read_input_tokens"] = int_parts["cached_prompt_tokens"]
|
||||||
|
if "uncached_prompt_tokens" in int_parts:
|
||||||
|
result["uncached_prompt_tokens"] = int_parts["uncached_prompt_tokens"]
|
||||||
|
if cost_usd is not None:
|
||||||
|
result["cost_usd"] = cost_usd
|
||||||
|
if model:
|
||||||
|
result["model"] = model
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _get_header_case_insensitive(
|
def _get_header_case_insensitive(
|
||||||
headers: Mapping[str, str], header_name: str
|
headers: Mapping[str, str], header_name: str
|
||||||
@@ -191,7 +274,9 @@ def _resolve_ehbp_target_url(
|
|||||||
otherwise the header is ignored so callers cannot redirect other providers
|
otherwise the header is ignored so callers cannot redirect other providers
|
||||||
or leak upstream API keys.
|
or leak upstream API keys.
|
||||||
"""
|
"""
|
||||||
override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
|
override_header = (
|
||||||
|
profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER
|
||||||
|
)
|
||||||
if not override_header:
|
if not override_header:
|
||||||
return target_url
|
return target_url
|
||||||
enclave_url = _get_header_case_insensitive(headers, override_header)
|
enclave_url = _get_header_case_insensitive(headers, override_header)
|
||||||
@@ -274,6 +359,11 @@ def _build_cost_info(
|
|||||||
output_tokens: int = 0,
|
output_tokens: int = 0,
|
||||||
input_msats: int = 0,
|
input_msats: int = 0,
|
||||||
output_msats: int = 0,
|
output_msats: int = 0,
|
||||||
|
cache_read_input_tokens: int = 0,
|
||||||
|
cache_creation_input_tokens: int = 0,
|
||||||
|
cache_read_msats: int = 0,
|
||||||
|
cache_creation_msats: int = 0,
|
||||||
|
total_usd: float = 0.0,
|
||||||
actual_model: str | None = None,
|
actual_model: str | None = None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""Build a cost-info dict with token counts and per-token-type costs.
|
"""Build a cost-info dict with token counts and per-token-type costs.
|
||||||
@@ -282,22 +372,25 @@ def _build_cost_info(
|
|||||||
one), it is included in the returned dict so callers can use it for billing
|
one), it is included in the returned dict so callers can use it for billing
|
||||||
finalization and logging.
|
finalization and logging.
|
||||||
"""
|
"""
|
||||||
result: dict[str, int | str | None] = {
|
result: dict[str, int | float | str | None] = {
|
||||||
"total_msats": total_msats,
|
"total_msats": total_msats,
|
||||||
"input_tokens": input_tokens,
|
"input_tokens": input_tokens,
|
||||||
"output_tokens": output_tokens,
|
"output_tokens": output_tokens,
|
||||||
"total_tokens": input_tokens + output_tokens,
|
"total_tokens": input_tokens + output_tokens,
|
||||||
"input_msats": input_msats,
|
"input_msats": input_msats,
|
||||||
"output_msats": output_msats,
|
"output_msats": output_msats,
|
||||||
|
"cache_read_input_tokens": cache_read_input_tokens,
|
||||||
|
"cache_creation_input_tokens": cache_creation_input_tokens,
|
||||||
|
"cache_read_msats": cache_read_msats,
|
||||||
|
"cache_creation_msats": cache_creation_msats,
|
||||||
|
"total_usd": total_usd,
|
||||||
}
|
}
|
||||||
if actual_model:
|
if actual_model:
|
||||||
result["actual_model"] = actual_model
|
result["actual_model"] = actual_model
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _inject_cost_response_headers(
|
def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None:
|
||||||
headers: dict[str, str], cost_info: dict
|
|
||||||
) -> None:
|
|
||||||
"""Add per-request cost headers to an EHBP response.
|
"""Add per-request cost headers to an EHBP response.
|
||||||
|
|
||||||
Since EHBP response bodies are opaque encrypted blobs, cost cannot be
|
Since EHBP response bodies are opaque encrypted blobs, cost cannot be
|
||||||
@@ -305,8 +398,14 @@ def _inject_cost_response_headers(
|
|||||||
the client/Tinfoil SDK can read without decrypting.
|
the client/Tinfoil SDK can read without decrypting.
|
||||||
"""
|
"""
|
||||||
headers["X-Routstr-Cost-Msats"] = str(cost_info["total_msats"])
|
headers["X-Routstr-Cost-Msats"] = str(cost_info["total_msats"])
|
||||||
|
if "computed_msats" in cost_info:
|
||||||
|
headers["X-Routstr-Computed-Cost-Msats"] = str(cost_info["computed_msats"])
|
||||||
headers["X-Routstr-Input-Cost-Msats"] = str(cost_info["input_msats"])
|
headers["X-Routstr-Input-Cost-Msats"] = str(cost_info["input_msats"])
|
||||||
headers["X-Routstr-Output-Cost-Msats"] = str(cost_info["output_msats"])
|
headers["X-Routstr-Output-Cost-Msats"] = str(cost_info["output_msats"])
|
||||||
|
headers["X-Routstr-Cache-Read-Msats"] = str(cost_info.get("cache_read_msats", 0))
|
||||||
|
headers["X-Routstr-Cache-Creation-Msats"] = str(
|
||||||
|
cost_info.get("cache_creation_msats", 0)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def _compute_ehbp_actual_cost(
|
async def _compute_ehbp_actual_cost(
|
||||||
@@ -316,10 +415,10 @@ async def _compute_ehbp_actual_cost(
|
|||||||
) -> dict:
|
) -> dict:
|
||||||
"""Compute the actual cost in msats from Tinfoil usage metrics.
|
"""Compute the actual cost in msats from Tinfoil usage metrics.
|
||||||
|
|
||||||
Falls back to ``max_cost_for_model`` when usage is absent (streaming) or
|
When usage is present, the result is clamped to ``[min_request_msat,
|
||||||
cannot be priced. The result is clamped to ``[min_request_msat,
|
max_cost_for_model]``. Missing or unpriceable usage returns zero: encrypted
|
||||||
max_cost_for_model]`` so the refund never exceeds the reservation and is
|
EHBP bodies cannot be estimated locally, and the authorization ceiling is
|
||||||
never zero.
|
not evidence of consumption.
|
||||||
|
|
||||||
When the usage-metrics header includes ``model=<name>`` and it differs
|
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
|
from ``model_obj.id``, the actual served model's pricing is used for the
|
||||||
@@ -332,7 +431,7 @@ async def _compute_ehbp_actual_cost(
|
|||||||
"""
|
"""
|
||||||
usage_dict = parse_tinfoil_usage_metrics(usage_header)
|
usage_dict = parse_tinfoil_usage_metrics(usage_header)
|
||||||
if usage_dict is None:
|
if usage_dict is None:
|
||||||
return _build_cost_info(max_cost_for_model)
|
return _build_cost_info(0)
|
||||||
|
|
||||||
# The enclave may serve a different model than the one requested (e.g.
|
# The enclave may serve a different model than the one requested (e.g.
|
||||||
# due to failover). The usage-metrics header's ``model=<name>`` carries
|
# due to failover). The usage-metrics header's ``model=<name>`` carries
|
||||||
@@ -344,6 +443,14 @@ async def _compute_ehbp_actual_cost(
|
|||||||
# look up the actual model's pricing.
|
# look up the actual model's pricing.
|
||||||
actual_model: str | None = usage_dict.pop("model", None) # type: ignore[arg-type]
|
actual_model: str | None = usage_dict.pop("model", None) # type: ignore[arg-type]
|
||||||
pricing_model_id = model_obj.id
|
pricing_model_id = model_obj.id
|
||||||
|
# Bill the model we actually routed to. Passing only the model *string*
|
||||||
|
# to calculate_cost makes it re-derive pricing from the global alias map,
|
||||||
|
# which resolves the id to the best-ranked candidate — not the serving
|
||||||
|
# one. Tinfoil's catalog id (e.g. ``deepseek-v4-1-flash``) is also a
|
||||||
|
# cross-provider alias, and that cheaper candidate has no cache rate, so
|
||||||
|
# the cache discount silently disappeared (and the request was
|
||||||
|
# undercharged). Hand calculate_cost the identity it cannot reconstruct.
|
||||||
|
pricing_model_obj: Model = model_obj
|
||||||
expected_upstream_model = model_obj.forwarded_model_id or model_obj.id
|
expected_upstream_model = model_obj.forwarded_model_id or model_obj.id
|
||||||
expected_identity = _normalize_upstream_model_id(expected_upstream_model)
|
expected_identity = _normalize_upstream_model_id(expected_upstream_model)
|
||||||
served_identity = _normalize_upstream_model_id(actual_model)
|
served_identity = _normalize_upstream_model_id(actual_model)
|
||||||
@@ -359,6 +466,23 @@ async def _compute_ehbp_actual_cost(
|
|||||||
# the global model map. The resolved object can belong to a different
|
# the global model map. The resolved object can belong to a different
|
||||||
# provider and therefore have a different client-facing ``id`` while
|
# provider and therefore have a different client-facing ``id`` while
|
||||||
# still representing the same upstream model.
|
# still representing the same upstream model.
|
||||||
|
#
|
||||||
|
# The enclave reports the *bare* upstream id, but the routstr model is
|
||||||
|
# namespaced ``tinfoil-`` (and the SDK strips that prefix for the
|
||||||
|
# encrypted body). Resolve the served id within the same namespace
|
||||||
|
# first: a same-model report then maps back onto the requested Tinfoil
|
||||||
|
# model, and a genuine failover lands on the actually-served Tinfoil
|
||||||
|
# model — instead of the cheaper cross-provider model the bare id
|
||||||
|
# would resolve to in the global map.
|
||||||
|
namespaced_served = actual_model
|
||||||
|
if (
|
||||||
|
expected_upstream_model.startswith(TINFOIL_MODEL_PREFIX)
|
||||||
|
and not actual_model.startswith(TINFOIL_MODEL_PREFIX)
|
||||||
|
):
|
||||||
|
namespaced_served = TINFOIL_MODEL_PREFIX + actual_model
|
||||||
|
|
||||||
|
actual_model_obj = get_model_instance(namespaced_served)
|
||||||
|
if actual_model_obj is None and namespaced_served != actual_model:
|
||||||
actual_model_obj = get_model_instance(actual_model)
|
actual_model_obj = get_model_instance(actual_model)
|
||||||
if actual_model_obj is None:
|
if actual_model_obj is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -375,9 +499,7 @@ async def _compute_ehbp_actual_cost(
|
|||||||
resolved_upstream_model = (
|
resolved_upstream_model = (
|
||||||
actual_model_obj.forwarded_model_id or actual_model_obj.id
|
actual_model_obj.forwarded_model_id or actual_model_obj.id
|
||||||
)
|
)
|
||||||
resolved_identity = _normalize_upstream_model_id(
|
resolved_identity = _normalize_upstream_model_id(resolved_upstream_model)
|
||||||
resolved_upstream_model
|
|
||||||
)
|
|
||||||
if resolved_identity != expected_identity:
|
if resolved_identity != expected_identity:
|
||||||
logger.info(
|
logger.info(
|
||||||
"EHBP served model differs from requested, using actual "
|
"EHBP served model differs from requested, using actual "
|
||||||
@@ -390,6 +512,7 @@ async def _compute_ehbp_actual_cost(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
pricing_model_id = actual_model_obj.id
|
pricing_model_id = actual_model_obj.id
|
||||||
|
pricing_model_obj = actual_model_obj
|
||||||
else:
|
else:
|
||||||
# A different registry/client alias resolved to the same
|
# A different registry/client alias resolved to the same
|
||||||
# upstream model; retain the requested model's pricing.
|
# upstream model; retain the requested model's pricing.
|
||||||
@@ -402,22 +525,23 @@ async def _compute_ehbp_actual_cost(
|
|||||||
cost = await calculate_cost(
|
cost = await calculate_cost(
|
||||||
{"model": pricing_model_id, "usage": usage_dict},
|
{"model": pricing_model_id, "usage": usage_dict},
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
|
pricing_model_obj,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"EHBP usage cost calculation failed, falling back to max cost",
|
"EHBP usage cost calculation failed; releasing instead of charging max cost",
|
||||||
extra={
|
extra={
|
||||||
"model": pricing_model_id,
|
"model": pricing_model_id,
|
||||||
"error": str(e),
|
"error": str(e),
|
||||||
"usage": usage_dict,
|
"usage": usage_dict,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
|
return _build_cost_info(0, actual_model=actual_model)
|
||||||
|
|
||||||
if isinstance(cost, MaxCostData):
|
if isinstance(cost, MaxCostData):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"EHBP calculate_cost returned MaxCostData (no model pricing), "
|
"EHBP calculate_cost returned MaxCostData (no usable pricing); "
|
||||||
"falling back to max cost",
|
"releasing instead of charging max cost",
|
||||||
extra={
|
extra={
|
||||||
"model": pricing_model_id,
|
"model": pricing_model_id,
|
||||||
"max_cost_for_model": max_cost_for_model,
|
"max_cost_for_model": max_cost_for_model,
|
||||||
@@ -425,7 +549,7 @@ async def _compute_ehbp_actual_cost(
|
|||||||
"cost_total_msats": cost.total_msats,
|
"cost_total_msats": cost.total_msats,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
|
return _build_cost_info(0, actual_model=actual_model)
|
||||||
if isinstance(cost, CostData):
|
if isinstance(cost, CostData):
|
||||||
actual = max(int(cost.total_msats), int(settings.min_request_msat))
|
actual = max(int(cost.total_msats), int(settings.min_request_msat))
|
||||||
clamped = min(actual, max_cost_for_model)
|
clamped = min(actual, max_cost_for_model)
|
||||||
@@ -445,17 +569,22 @@ async def _compute_ehbp_actual_cost(
|
|||||||
output_tokens=cost.output_tokens,
|
output_tokens=cost.output_tokens,
|
||||||
input_msats=cost.input_msats,
|
input_msats=cost.input_msats,
|
||||||
output_msats=cost.output_msats,
|
output_msats=cost.output_msats,
|
||||||
|
cache_read_input_tokens=cost.cache_read_input_tokens,
|
||||||
|
cache_creation_input_tokens=cost.cache_creation_input_tokens,
|
||||||
|
cache_read_msats=cost.cache_read_msats,
|
||||||
|
cache_creation_msats=cost.cache_creation_msats,
|
||||||
|
total_usd=cost.total_usd,
|
||||||
actual_model=actual_model,
|
actual_model=actual_model,
|
||||||
)
|
)
|
||||||
# CostDataError
|
# CostDataError
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"EHBP usage cost calculation error, falling back to max cost",
|
"EHBP usage cost calculation error; releasing instead of charging max cost",
|
||||||
extra={
|
extra={
|
||||||
"model": pricing_model_id,
|
"model": pricing_model_id,
|
||||||
"error": getattr(cost, "message", str(cost)),
|
"error": getattr(cost, "message", str(cost)),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return _build_cost_info(max_cost_for_model, actual_model=actual_model)
|
return _build_cost_info(0, actual_model=actual_model)
|
||||||
|
|
||||||
|
|
||||||
def _extract_usage_from_response(
|
def _extract_usage_from_response(
|
||||||
@@ -500,6 +629,18 @@ class EHBPForwardingTarget:
|
|||||||
profile: ConfidentialInferenceProfile | None = None
|
profile: ConfidentialInferenceProfile | None = None
|
||||||
|
|
||||||
|
|
||||||
|
async def _release_failed_ehbp_charge(
|
||||||
|
reservation: ReservationSnapshot, session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
if await release_reservation(reservation, session, reservation.reserved_msats):
|
||||||
|
return
|
||||||
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
logger.critical(
|
||||||
|
"Failed to release EHBP reservation after rejected charge",
|
||||||
|
extra={"reservation_id": reservation.release_id},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def finalize_ehbp_actual_cost_payment(
|
async def finalize_ehbp_actual_cost_payment(
|
||||||
key: ApiKey,
|
key: ApiKey,
|
||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
@@ -507,61 +648,27 @@ async def finalize_ehbp_actual_cost_payment(
|
|||||||
model_id: str,
|
model_id: str,
|
||||||
cost_info: dict,
|
cost_info: dict,
|
||||||
reservation_snapshot: ReservationSnapshot | None = None,
|
reservation_snapshot: ReservationSnapshot | None = None,
|
||||||
) -> None:
|
) -> int:
|
||||||
"""Finalize an EHBP bearer request using clamped provider usage metrics."""
|
"""Finalize an EHBP bearer request using clamped provider usage metrics."""
|
||||||
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
||||||
await _validate_reservation_snapshot(key, reservation, session)
|
await _validate_reservation_snapshot(key, reservation, session)
|
||||||
if not await _claim_reservation_for_charge(reservation, session):
|
if not await _claim_reservation_for_charge(reservation, session):
|
||||||
return
|
return 0
|
||||||
reserved_cost_for_model = reservation.reserved_msats
|
reserved_cost_for_model = reservation.reserved_msats
|
||||||
billing_key = await get_billing_key(key, session)
|
|
||||||
key_hash = key.hashed_key
|
key_hash = key.hashed_key
|
||||||
billing_key_hash = billing_key.hashed_key
|
billing_key_hash = key_hash
|
||||||
total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model)))
|
total_cost_msats = max(
|
||||||
|
0, int(cost_info.get("total_msats", reserved_cost_for_model))
|
||||||
|
)
|
||||||
now = int(time.time())
|
now = int(time.time())
|
||||||
|
|
||||||
safe_reserved = case(
|
charged = await _charge_reservation_rows(
|
||||||
(
|
session,
|
||||||
col(ApiKey.reserved_balance) >= reserved_cost_for_model,
|
billing_key_hash=billing_key_hash,
|
||||||
col(ApiKey.reserved_balance) - reserved_cost_for_model,
|
reserved_msats=reserved_cost_for_model,
|
||||||
),
|
charge_msats=total_cost_msats,
|
||||||
else_=0,
|
|
||||||
)
|
)
|
||||||
cleared_reserved_at = case(
|
if not charged:
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) - reserved_cost_for_model > 0,
|
|
||||||
col(ApiKey.reserved_at),
|
|
||||||
),
|
|
||||||
else_=None,
|
|
||||||
)
|
|
||||||
|
|
||||||
stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=safe_reserved,
|
|
||||||
reserved_at=cleared_reserved_at,
|
|
||||||
balance=col(ApiKey.balance) - total_cost_msats,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
child_result = None
|
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
|
||||||
child_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=safe_reserved,
|
|
||||||
reserved_at=cleared_reserved_at,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
|
|
||||||
await session.rollback()
|
|
||||||
logger.error(
|
logger.error(
|
||||||
"Failed to finalize EHBP usage-based payment",
|
"Failed to finalize EHBP usage-based payment",
|
||||||
extra={
|
extra={
|
||||||
@@ -570,15 +677,13 @@ async def finalize_ehbp_actual_cost_payment(
|
|||||||
"model": model_id,
|
"model": model_id,
|
||||||
"reserved_cost_for_model": reserved_cost_for_model,
|
"reserved_cost_for_model": reserved_cost_for_model,
|
||||||
"total_cost_msats": total_cost_msats,
|
"total_cost_msats": total_cost_msats,
|
||||||
"parent_rowcount": result.rowcount,
|
|
||||||
"child_rowcount": getattr(child_result, "rowcount", None),
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return
|
await _release_failed_ehbp_charge(reservation, session)
|
||||||
|
return 0
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
await session.refresh(billing_key)
|
await _stop_reservation_heartbeat(reservation.release_id)
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
|
||||||
await session.refresh(key)
|
await session.refresh(key)
|
||||||
|
|
||||||
if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0:
|
if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0:
|
||||||
@@ -596,19 +701,30 @@ async def finalize_ehbp_actual_cost_payment(
|
|||||||
extra={
|
extra={
|
||||||
"event": "finalize",
|
"event": "finalize",
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_hash": key.hashed_key[:8] + "...",
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
"billing_key_hash": key.hashed_key[:8] + "...",
|
||||||
"model": model_id,
|
"model": model_id,
|
||||||
"cost_reserved": reserved_cost_for_model,
|
"cost_reserved": reserved_cost_for_model,
|
||||||
"cost_charged": total_cost_msats,
|
"cost_charged": total_cost_msats,
|
||||||
"input_tokens": cost_info.get("input_tokens", 0),
|
"input_tokens": cost_info.get("input_tokens", 0),
|
||||||
"output_tokens": cost_info.get("output_tokens", 0),
|
"output_tokens": cost_info.get("output_tokens", 0),
|
||||||
"balance": billing_key.balance,
|
# Cache splits are only knowable when the enclave reports
|
||||||
"reserved_balance": billing_key.reserved_balance,
|
# ``cached_prompt_tokens``; absent that they are a measured zero on
|
||||||
"total_spent": billing_key.total_spent,
|
# the token counts the provider did report (not an unknown), so the
|
||||||
|
# event key set stays stable for usage-analytics consumers.
|
||||||
|
"cache_read_input_tokens": cost_info.get("cache_read_input_tokens", 0),
|
||||||
|
"cache_creation_input_tokens": cost_info.get(
|
||||||
|
"cache_creation_input_tokens", 0
|
||||||
|
),
|
||||||
|
"cache_read_msats": cost_info.get("cache_read_msats", 0),
|
||||||
|
"cache_creation_msats": cost_info.get("cache_creation_msats", 0),
|
||||||
|
"balance": key.balance,
|
||||||
|
"reserved_balance": key.reserved_balance,
|
||||||
|
"total_spent": key.total_spent,
|
||||||
"finalize_type": "ehbp_usage",
|
"finalize_type": "ehbp_usage",
|
||||||
"finalized_at": now,
|
"finalized_at": now,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
return total_cost_msats
|
||||||
|
|
||||||
|
|
||||||
async def finalize_ehbp_max_cost_payment(
|
async def finalize_ehbp_max_cost_payment(
|
||||||
@@ -617,128 +733,26 @@ async def finalize_ehbp_max_cost_payment(
|
|||||||
max_cost_for_model: int,
|
max_cost_for_model: int,
|
||||||
model_id: str,
|
model_id: str,
|
||||||
reservation_snapshot: ReservationSnapshot | None = None,
|
reservation_snapshot: ReservationSnapshot | None = None,
|
||||||
) -> None:
|
) -> int:
|
||||||
"""Finalize an EHBP bearer request by charging the reserved max cost.
|
"""Release an unmeasured EHBP request without charging its reservation.
|
||||||
|
|
||||||
EHBP responses are encrypted, so Routstr cannot inspect token usage. Unlike
|
The legacy name is retained for compatibility with internal callers. EHBP
|
||||||
normal completion handlers, this intentionally charges the pre-reserved max
|
responses are encrypted, so no local estimate is possible when the trusted
|
||||||
cost and releases the reservation.
|
usage header/trailer is absent.
|
||||||
"""
|
"""
|
||||||
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
reservation = reservation_snapshot or await get_reservation_snapshot(key, session)
|
||||||
await _validate_reservation_snapshot(key, reservation, session)
|
await _validate_reservation_snapshot(key, reservation, session)
|
||||||
if not await _claim_reservation_for_charge(reservation, session):
|
key_log_hash = key.hashed_key[:8] + "..."
|
||||||
return
|
await release_reservation(reservation, session, reservation.reserved_msats)
|
||||||
max_cost_for_model = reservation.reserved_msats
|
logger.warning(
|
||||||
billing_key = await get_billing_key(key, session)
|
"Released unmeasured EHBP reservation without charging max cost",
|
||||||
key_hash = key.hashed_key
|
|
||||||
billing_key_hash = billing_key.hashed_key
|
|
||||||
total_cost_msats = max(0, int(max_cost_for_model))
|
|
||||||
now = int(time.time())
|
|
||||||
|
|
||||||
cleared_reserved_at = case(
|
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) - max_cost_for_model > 0,
|
|
||||||
col(ApiKey.reserved_at),
|
|
||||||
),
|
|
||||||
else_=None,
|
|
||||||
)
|
|
||||||
safe_reserved = case(
|
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) >= max_cost_for_model,
|
|
||||||
col(ApiKey.reserved_balance) - max_cost_for_model,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
|
||||||
|
|
||||||
stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=safe_reserved,
|
|
||||||
reserved_at=cleared_reserved_at,
|
|
||||||
balance=col(ApiKey.balance) - total_cost_msats,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
|
||||||
|
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
|
||||||
child_safe_reserved = case(
|
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) >= max_cost_for_model,
|
|
||||||
col(ApiKey.reserved_balance) - max_cost_for_model,
|
|
||||||
),
|
|
||||||
else_=0,
|
|
||||||
)
|
|
||||||
child_cleared_reserved_at = case(
|
|
||||||
(
|
|
||||||
col(ApiKey.reserved_balance) - max_cost_for_model > 0,
|
|
||||||
col(ApiKey.reserved_at),
|
|
||||||
),
|
|
||||||
else_=None,
|
|
||||||
)
|
|
||||||
child_stmt = (
|
|
||||||
update(ApiKey)
|
|
||||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
|
||||||
.values(
|
|
||||||
reserved_balance=child_safe_reserved,
|
|
||||||
reserved_at=child_cleared_reserved_at,
|
|
||||||
total_spent=col(ApiKey.total_spent) + total_cost_msats,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
child_result = await session.exec(child_stmt) # type: ignore[call-overload]
|
|
||||||
else:
|
|
||||||
child_result = None
|
|
||||||
|
|
||||||
if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0):
|
|
||||||
await session.rollback()
|
|
||||||
logger.error(
|
|
||||||
"Failed to finalize EHBP max-cost payment",
|
|
||||||
extra={
|
extra={
|
||||||
"key_hash": key_hash[:8] + "...",
|
"key_hash": key_log_hash,
|
||||||
"billing_key_hash": billing_key_hash[:8] + "...",
|
|
||||||
"model": model_id,
|
"model": model_id,
|
||||||
"max_cost_for_model": max_cost_for_model,
|
"max_cost_for_model": max_cost_for_model,
|
||||||
"parent_rowcount": result.rowcount,
|
|
||||||
"child_rowcount": getattr(child_result, "rowcount", None),
|
|
||||||
},
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
await session.refresh(billing_key)
|
|
||||||
if billing_key.hashed_key != key.hashed_key:
|
|
||||||
await session.refresh(key)
|
|
||||||
|
|
||||||
if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0:
|
|
||||||
fee_msats = math.ceil(total_cost_msats * ROUTSTR_FEE_PERCENT / 100)
|
|
||||||
try:
|
|
||||||
await accumulate_routstr_fee(session, fee_msats)
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(
|
|
||||||
"Failed to accumulate Routstr fee for EHBP request",
|
|
||||||
extra={"error": str(e), "fee_msats": fee_msats},
|
|
||||||
)
|
|
||||||
|
|
||||||
payments_logger.info(
|
|
||||||
"FINALIZE",
|
|
||||||
extra={
|
|
||||||
"event": "finalize",
|
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
|
||||||
"billing_key_hash": billing_key.hashed_key[:8] + "...",
|
|
||||||
"model": model_id,
|
|
||||||
"cost_reserved": max_cost_for_model,
|
|
||||||
"cost_charged": total_cost_msats,
|
|
||||||
"input_tokens": 0,
|
|
||||||
"output_tokens": 0,
|
|
||||||
"balance": billing_key.balance,
|
|
||||||
"reserved_balance": billing_key.reserved_balance,
|
|
||||||
"total_spent": billing_key.total_spent,
|
|
||||||
"finalize_type": "ehbp_max_cost",
|
|
||||||
"finalized_at": now,
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
async def send_cashu_refund(
|
async def send_cashu_refund(
|
||||||
@@ -842,6 +856,25 @@ async def forward_ehbp_request(
|
|||||||
"body_preview": body_preview,
|
"body_preview": body_preview,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
# Key-config mismatch (stale client HPKE key): return the
|
||||||
|
# enclave's 422 problem+json directly so the SDK's
|
||||||
|
# KeyConfigMismatchError detection fires and triggers
|
||||||
|
# re-attestation. Wrapping it as application/json would
|
||||||
|
# destroy the signal and cause permanent failure.
|
||||||
|
if _is_ehbp_key_config_response(resp):
|
||||||
|
logger.warning(
|
||||||
|
"EHBP upstream %s returned key-config mismatch for model=%s, "
|
||||||
|
"passing through for client re-attestation",
|
||||||
|
provider_type,
|
||||||
|
model_obj.id,
|
||||||
|
extra={
|
||||||
|
"provider": provider_type,
|
||||||
|
"model": model_obj.id,
|
||||||
|
"path": path,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return _passthrough_key_config_response(resp)
|
||||||
|
|
||||||
raise UpstreamError(
|
raise UpstreamError(
|
||||||
f"EHBP upstream {provider_type} returned {resp.status_code} "
|
f"EHBP upstream {provider_type} returned {resp.status_code} "
|
||||||
f"for model {model_obj.id}: {body_preview[:200] or '<empty>'}",
|
f"for model {model_obj.id}: {body_preview[:200] or '<empty>'}",
|
||||||
@@ -893,10 +926,9 @@ async def forward_ehbp_request(
|
|||||||
cost_info = await _compute_ehbp_actual_cost(
|
cost_info = await _compute_ehbp_actual_cost(
|
||||||
usage_header, model_obj, max_cost_for_model
|
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
|
billing_model = cost_info.pop("actual_model", None) or model_obj.id
|
||||||
await finalize_ehbp_actual_cost_payment(
|
computed_msats = int(cost_info["total_msats"])
|
||||||
|
charged_msats = await finalize_ehbp_actual_cost_payment(
|
||||||
key,
|
key,
|
||||||
session,
|
session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
@@ -904,18 +936,25 @@ async def forward_ehbp_request(
|
|||||||
cost_info,
|
cost_info,
|
||||||
reservation_snapshot,
|
reservation_snapshot,
|
||||||
)
|
)
|
||||||
cost_data = {**cost_info, "total_usd": 0.0}
|
cost_data = {
|
||||||
|
**cost_info,
|
||||||
|
"total_msats": charged_msats,
|
||||||
|
"charged_msats": charged_msats,
|
||||||
|
"total_usd": cost_info.get("total_usd", 0.0),
|
||||||
|
}
|
||||||
|
if computed_msats != charged_msats:
|
||||||
|
cost_data["computed_msats"] = computed_msats
|
||||||
else:
|
else:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"EHBP usage metrics not found in headers or trailers, "
|
"EHBP usage metrics not found in headers or trailers; "
|
||||||
"falling back to max-cost billing",
|
"releasing instead of charging the authorization ceiling",
|
||||||
extra={
|
extra={
|
||||||
"model": model_obj.id,
|
"model": model_obj.id,
|
||||||
"provider": provider_type,
|
"provider": provider_type,
|
||||||
"key_hash": key.hashed_key[:8] + "...",
|
"key_hash": key.hashed_key[:8] + "...",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
await finalize_ehbp_max_cost_payment(
|
charged_msats = await finalize_ehbp_max_cost_payment(
|
||||||
key,
|
key,
|
||||||
session,
|
session,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
@@ -923,14 +962,15 @@ async def forward_ehbp_request(
|
|||||||
reservation_snapshot,
|
reservation_snapshot,
|
||||||
)
|
)
|
||||||
cost_data = {
|
cost_data = {
|
||||||
"total_msats": max_cost_for_model,
|
"total_msats": charged_msats,
|
||||||
|
"charged_msats": charged_msats,
|
||||||
"total_usd": 0.0,
|
"total_usd": 0.0,
|
||||||
"input_tokens": 0,
|
"input_tokens": 0,
|
||||||
"output_tokens": 0,
|
"output_tokens": 0,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Build the cost_info dict from what adjust_payment_for_tokens returned
|
# Build the cost_info dict from measured usage or the unmeasured-release
|
||||||
# or from the max-cost fallback. Fields match CostData/MaxCostData.dict().
|
# fallback. Fields match CostData/MaxCostData.dict().
|
||||||
cost_info = {
|
cost_info = {
|
||||||
"total_msats": cost_data.get("total_msats", max_cost_for_model),
|
"total_msats": cost_data.get("total_msats", max_cost_for_model),
|
||||||
"input_tokens": cost_data.get("input_tokens", 0),
|
"input_tokens": cost_data.get("input_tokens", 0),
|
||||||
@@ -939,7 +979,15 @@ async def forward_ehbp_request(
|
|||||||
+ cost_data.get("output_tokens", 0),
|
+ cost_data.get("output_tokens", 0),
|
||||||
"input_msats": cost_data.get("input_msats", 0),
|
"input_msats": cost_data.get("input_msats", 0),
|
||||||
"output_msats": cost_data.get("output_msats", 0),
|
"output_msats": cost_data.get("output_msats", 0),
|
||||||
|
"cache_read_input_tokens": cost_data.get("cache_read_input_tokens", 0),
|
||||||
|
"cache_creation_input_tokens": cost_data.get(
|
||||||
|
"cache_creation_input_tokens", 0
|
||||||
|
),
|
||||||
|
"cache_read_msats": cost_data.get("cache_read_msats", 0),
|
||||||
|
"cache_creation_msats": cost_data.get("cache_creation_msats", 0),
|
||||||
}
|
}
|
||||||
|
if "computed_msats" in cost_data:
|
||||||
|
cost_info["computed_msats"] = cost_data["computed_msats"]
|
||||||
cost_usd = cost_data.get("total_usd", 0.0)
|
cost_usd = cost_data.get("total_usd", 0.0)
|
||||||
|
|
||||||
# Build response headers, filtering out hop-by-hop headers
|
# Build response headers, filtering out hop-by-hop headers
|
||||||
@@ -1028,7 +1076,9 @@ async def forward_ehbp_x_cashu_request(
|
|||||||
target_url = _resolve_ehbp_target_url(
|
target_url = _resolve_ehbp_target_url(
|
||||||
target.url, path, headers, provider_type, profile
|
target.url, path, headers, provider_type, profile
|
||||||
)
|
)
|
||||||
upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile)
|
upstream_headers = _prepare_ehbp_upstream_headers(
|
||||||
|
headers, target.headers, profile
|
||||||
|
)
|
||||||
request_body = await request.body()
|
request_body = await request.body()
|
||||||
|
|
||||||
# Merge query params into the target URL
|
# Merge query params into the target URL
|
||||||
@@ -1047,6 +1097,29 @@ async def forward_ehbp_x_cashu_request(
|
|||||||
)
|
)
|
||||||
|
|
||||||
if resp.status_code != 200:
|
if resp.status_code != 200:
|
||||||
|
# Key-config mismatch (stale client HPKE key): refund the
|
||||||
|
# full token and pass the enclave's 422 problem+json through
|
||||||
|
# so the SDK's KeyConfigMismatchError detection fires.
|
||||||
|
if _is_ehbp_key_config_response(resp):
|
||||||
|
logger.warning(
|
||||||
|
"EHBP upstream %s returned key-config mismatch for "
|
||||||
|
"model=%s, refunding and passing through",
|
||||||
|
provider_type,
|
||||||
|
model_obj.id,
|
||||||
|
extra={
|
||||||
|
"provider": provider_type,
|
||||||
|
"model": model_obj.id,
|
||||||
|
"path": path,
|
||||||
|
"refunded_amount": amount,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
refund_token = await send_cashu_refund(
|
||||||
|
amount, unit, mint, request_id
|
||||||
|
)
|
||||||
|
passthrough = _passthrough_key_config_response(resp)
|
||||||
|
passthrough.headers["X-Cashu"] = refund_token
|
||||||
|
return passthrough
|
||||||
|
|
||||||
refund_token = await send_cashu_refund(amount, unit, mint, request_id)
|
refund_token = await send_cashu_refund(amount, unit, mint, request_id)
|
||||||
error_response = Response(
|
error_response = Response(
|
||||||
content=json.dumps(
|
content=json.dumps(
|
||||||
@@ -1076,9 +1149,7 @@ async def forward_ehbp_x_cashu_request(
|
|||||||
usage_source = (
|
usage_source = (
|
||||||
"header"
|
"header"
|
||||||
if usage_header_name
|
if usage_header_name
|
||||||
and any(
|
and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers)
|
||||||
k.lower() == usage_header_name.lower() for k, _ in resp.headers
|
|
||||||
)
|
|
||||||
else ("trailer" if usage_header else "none")
|
else ("trailer" if usage_header else "none")
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1153,6 +1224,46 @@ async def forward_ehbp_x_cashu_request(
|
|||||||
except Exception:
|
except Exception:
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
except EhbpTimeoutError as e:
|
||||||
|
logger.warning(
|
||||||
|
"EHBP X-Cashu upstream timed out",
|
||||||
|
extra={
|
||||||
|
"error": str(e),
|
||||||
|
"path": path,
|
||||||
|
"method": request.method,
|
||||||
|
"redeemed": redeemed,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
if redeemed and amount > 0:
|
||||||
|
try:
|
||||||
|
refund_token = await send_cashu_refund(amount, unit, mint, request_id)
|
||||||
|
error_response = create_error_response(
|
||||||
|
"upstream_timeout",
|
||||||
|
str(e),
|
||||||
|
504,
|
||||||
|
request=request,
|
||||||
|
code="UPSTREAM_TIMEOUT",
|
||||||
|
)
|
||||||
|
error_response.headers["X-Cashu"] = refund_token
|
||||||
|
return error_response
|
||||||
|
except Exception as refund_error:
|
||||||
|
logger.error(
|
||||||
|
"Failed to refund EHBP X-Cashu token after timeout",
|
||||||
|
extra={
|
||||||
|
"error": str(refund_error),
|
||||||
|
"original_error": str(e),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
return create_error_response(
|
||||||
|
"upstream_timeout",
|
||||||
|
str(e),
|
||||||
|
504,
|
||||||
|
request=request,
|
||||||
|
code="UPSTREAM_TIMEOUT",
|
||||||
|
)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
error_message = str(e)
|
error_message = str(e)
|
||||||
logger.error(
|
logger.error(
|
||||||
@@ -1186,6 +1297,30 @@ async def forward_ehbp_x_cashu_request(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if not redeemed:
|
||||||
|
classified = classify_redemption_error(e)
|
||||||
|
if classified is not None:
|
||||||
|
error_type, status_code, message, error_code = classified
|
||||||
|
# Never re-offer a spent/consumed token.
|
||||||
|
echo_token = None if error_code in SPENT_TOKEN_CODES else x_cashu_token
|
||||||
|
return create_error_response(
|
||||||
|
error_type,
|
||||||
|
message,
|
||||||
|
status_code,
|
||||||
|
request=request,
|
||||||
|
token=echo_token,
|
||||||
|
code=error_code,
|
||||||
|
)
|
||||||
|
# Raw exception text may contain the attacker-supplied mint URL.
|
||||||
|
return create_error_response(
|
||||||
|
"api_error",
|
||||||
|
"Internal error during token redemption",
|
||||||
|
500,
|
||||||
|
request=request,
|
||||||
|
token=x_cashu_token,
|
||||||
|
code="internal_error",
|
||||||
|
)
|
||||||
|
|
||||||
if "already spent" in error_message.lower():
|
if "already spent" in error_message.lower():
|
||||||
return create_error_response(
|
return create_error_response(
|
||||||
"token_already_spent",
|
"token_already_spent",
|
||||||
|
|||||||
@@ -32,6 +32,7 @@ from ..core.exceptions import UpstreamError
|
|||||||
from ..core.redaction import redact_org_ids
|
from ..core.redaction import redact_org_ids
|
||||||
from ..payment.models import Model
|
from ..payment.models import Model
|
||||||
from .rate_limit import classify_rate_limit
|
from .rate_limit import classify_rate_limit
|
||||||
|
from .reasoning_effort import adapt_messages_body_for_litellm
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -40,6 +41,10 @@ logger = get_logger(__name__)
|
|||||||
# unsupported params; these newer/extension fields get passed through
|
# unsupported params; these newer/extension fields get passed through
|
||||||
# verbatim and the upstream rejects them with a 400. Pop them here so the
|
# verbatim and the upstream rejects them with a 400. Pop them here so the
|
||||||
# request reaches the upstream cleanly.
|
# request reaches the upstream cleanly.
|
||||||
|
#
|
||||||
|
# Note: ``dispatch_anthropic_messages`` additionally enforces
|
||||||
|
# ``ALLOWED_MESSAGES_REQUEST_FIELDS``, which already excludes all of these.
|
||||||
|
# This tuple remains for ``gemini_messages``, which pops them explicitly.
|
||||||
ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = (
|
ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = (
|
||||||
"thinking",
|
"thinking",
|
||||||
"cache_control",
|
"cache_control",
|
||||||
@@ -51,6 +56,27 @@ ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = (
|
|||||||
"anthropic_beta",
|
"anthropic_beta",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Anthropic Messages API request fields forwarded from the client body
|
||||||
|
# into the upstream call. Only these are passed on; anything else is
|
||||||
|
# dropped so the forwarded request is deterministic and limited to the
|
||||||
|
# documented Messages surface.
|
||||||
|
ALLOWED_MESSAGES_REQUEST_FIELDS: frozenset[str] = frozenset(
|
||||||
|
{
|
||||||
|
"messages",
|
||||||
|
"max_tokens",
|
||||||
|
"system",
|
||||||
|
"temperature",
|
||||||
|
"top_p",
|
||||||
|
"top_k",
|
||||||
|
"stop_sequences",
|
||||||
|
"tools",
|
||||||
|
"tool_choice",
|
||||||
|
"metadata",
|
||||||
|
# OpenAI-shaped effort after thinking is lifted off Anthropic bodies.
|
||||||
|
"reasoning_effort",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def coerce_litellm_payload(payload: object) -> dict:
|
def coerce_litellm_payload(payload: object) -> dict:
|
||||||
"""Convert a litellm event into a plain dict.
|
"""Convert a litellm event into a plain dict.
|
||||||
@@ -108,9 +134,7 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]:
|
|||||||
return events, buffer
|
return events, buffer
|
||||||
|
|
||||||
|
|
||||||
def events_from_chunk(
|
def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]:
|
||||||
chunk: object, sse_buffer: bytes
|
|
||||||
) -> tuple[list[dict], bytes]:
|
|
||||||
"""Normalize a stream chunk into one or more event dicts.
|
"""Normalize a stream chunk into one or more event dicts.
|
||||||
|
|
||||||
``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE
|
``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE
|
||||||
@@ -201,9 +225,7 @@ async def aggregate_anthropic_events_to_message(
|
|||||||
raw_json = partial_json.pop(idx, None)
|
raw_json = partial_json.pop(idx, None)
|
||||||
if raw_json is not None and idx < len(blocks):
|
if raw_json is not None and idx < len(blocks):
|
||||||
try:
|
try:
|
||||||
blocks[idx]["input"] = (
|
blocks[idx]["input"] = json.loads(raw_json) if raw_json else {}
|
||||||
json.loads(raw_json) if raw_json else {}
|
|
||||||
)
|
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
blocks[idx]["input"] = raw_json
|
blocks[idx]["input"] = raw_json
|
||||||
elif etype == "message_delta":
|
elif etype == "message_delta":
|
||||||
@@ -445,9 +467,7 @@ async def dispatch_anthropic_messages(
|
|||||||
on bad input or upstream failure.
|
on bad input or upstream failure.
|
||||||
"""
|
"""
|
||||||
if not request_body:
|
if not request_body:
|
||||||
raise UpstreamError(
|
raise UpstreamError("Missing request body for /v1/messages", status_code=400)
|
||||||
"Missing request body for /v1/messages", status_code=400
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
body: dict = json.loads(request_body)
|
body: dict = json.loads(request_body)
|
||||||
@@ -466,15 +486,18 @@ async def dispatch_anthropic_messages(
|
|||||||
client_stream = bool(body.pop("stream", False))
|
client_stream = bool(body.pop("stream", False))
|
||||||
upstream_stream = True
|
upstream_stream = True
|
||||||
|
|
||||||
dropped: dict[str, Any] = {}
|
adapt_messages_body_for_litellm(body, model_obj)
|
||||||
for field in ANTHROPIC_ONLY_FIELDS:
|
|
||||||
if field in body:
|
# Forward only allowlisted Anthropic Messages request fields. Any
|
||||||
dropped[field] = body.pop(field)
|
# other client-supplied key is dropped so it cannot leak into the
|
||||||
|
# upstream request. See ALLOWED_MESSAGES_REQUEST_FIELDS.
|
||||||
|
dropped = sorted(set(body) - ALLOWED_MESSAGES_REQUEST_FIELDS)
|
||||||
if dropped:
|
if dropped:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Dropped anthropic-only fields before litellm dispatch",
|
"Dropped non-forwardable fields before litellm dispatch",
|
||||||
extra={"dropped_keys": sorted(dropped.keys())},
|
extra={"dropped_keys": dropped},
|
||||||
)
|
)
|
||||||
|
body = {k: v for k, v in body.items() if k in ALLOWED_MESSAGES_REQUEST_FIELDS}
|
||||||
|
|
||||||
# Convention: `model.id` is the canonical upstream model name;
|
# Convention: `model.id` is the canonical upstream model name;
|
||||||
# `forwarded_model_id` is the public alias the internal API exposes
|
# `forwarded_model_id` is the public alias the internal API exposes
|
||||||
|
|||||||
+191
-24
@@ -1,9 +1,6 @@
|
|||||||
"""Model-path discovery service.
|
"""Model-path discovery service.
|
||||||
|
|
||||||
Exposes every selectable upstream route a Routstr model is reachable through.
|
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
|
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
|
upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter
|
||||||
endpoint, its machine-readable tag::
|
endpoint, its machine-readable tag::
|
||||||
@@ -16,11 +13,12 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import ipaddress
|
import ipaddress
|
||||||
|
import json
|
||||||
import random
|
import random
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Any, Callable
|
from typing import TYPE_CHECKING, Any, Callable
|
||||||
from urllib.parse import urlencode, urlsplit
|
from urllib.parse import parse_qsl, urlencode, urlsplit
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from sqlalchemy.dialects.sqlite import insert
|
from sqlalchemy.dialects.sqlite import insert
|
||||||
@@ -60,10 +58,11 @@ ModelKey = tuple[str, int]
|
|||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class EndpointIdentity:
|
class EndpointIdentity:
|
||||||
"""Exact OpenRouter endpoint identity returned by ``/endpoints``."""
|
"""Exact OpenRouter endpoint and its provider-specific model metadata."""
|
||||||
|
|
||||||
tag: str
|
tag: str
|
||||||
provider_name: str | None
|
provider_name: str | None
|
||||||
|
model_metadata: dict[str, Any]
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -83,6 +82,7 @@ class DiscoveredPath:
|
|||||||
model_id: str
|
model_id: str
|
||||||
path: str
|
path: str
|
||||||
provider: ConfiguredProviderIdentity
|
provider: ConfiguredProviderIdentity
|
||||||
|
model_metadata: dict[str, Any]
|
||||||
endpoint_tag: str | None = None
|
endpoint_tag: str | None = None
|
||||||
endpoint_name: str | None = None
|
endpoint_name: str | None = None
|
||||||
|
|
||||||
@@ -132,21 +132,68 @@ def encode_model_path(
|
|||||||
return urlencode(components)
|
return urlencode(components)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ModelPathSelector:
|
||||||
|
"""Decoded client-supplied route selector."""
|
||||||
|
|
||||||
|
base_url: str
|
||||||
|
provider_id: int
|
||||||
|
model_id: str
|
||||||
|
endpoint_tag: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def decode_model_path(path: str) -> ModelPathSelector | None:
|
||||||
|
"""Inverse of ``encode_model_path``; ``None`` when the selector is malformed."""
|
||||||
|
try:
|
||||||
|
pairs = parse_qsl(
|
||||||
|
path,
|
||||||
|
keep_blank_values=True,
|
||||||
|
strict_parsing=True,
|
||||||
|
max_num_fields=4,
|
||||||
|
errors="strict",
|
||||||
|
)
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
params = dict(pairs)
|
||||||
|
if len(params) != len(pairs) or params.keys() - {
|
||||||
|
"url",
|
||||||
|
"provider-id",
|
||||||
|
"model-id",
|
||||||
|
"endpoint",
|
||||||
|
}:
|
||||||
|
return None
|
||||||
|
if any(not value.strip() for value in params.values()):
|
||||||
|
return None
|
||||||
|
base_url = params.get("url", "")
|
||||||
|
model_id = params.get("model-id", "")
|
||||||
|
raw_provider_id = params.get("provider-id", "")
|
||||||
|
if not base_url or not model_id or not raw_provider_id:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
provider_id = int(raw_provider_id)
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
if provider_id <= 0:
|
||||||
|
return None
|
||||||
|
return ModelPathSelector(
|
||||||
|
base_url=base_url,
|
||||||
|
provider_id=provider_id,
|
||||||
|
model_id=model_id,
|
||||||
|
endpoint_tag=params.get("endpoint") or None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _make_http_client() -> httpx.AsyncClient:
|
def _make_http_client() -> httpx.AsyncClient:
|
||||||
"""Client factory, separated so tests can substitute a mock transport."""
|
"""Client factory, separated so tests can substitute a mock transport."""
|
||||||
return httpx.AsyncClient()
|
return httpx.AsyncClient()
|
||||||
|
|
||||||
|
|
||||||
def is_openrouter_base_url(base_url: str | None) -> bool:
|
def is_openrouter_base_url(base_url: str | None) -> bool:
|
||||||
"""True when ``base_url`` points at OpenRouter.
|
"""Match OpenRouter itself, not compatible providers or lookalike hosts."""
|
||||||
|
try:
|
||||||
Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``:
|
return urlsplit(base_url or "").hostname == "openrouter.ai"
|
||||||
that predicate also returns True for native Anthropic (correct for
|
except ValueError:
|
||||||
cache-control, wrong for OpenRouter endpoint discovery). This one keys only
|
return False
|
||||||
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:
|
def exposed_model_id(model: object) -> str:
|
||||||
@@ -263,9 +310,14 @@ async def _fetch_openrouter_endpoint_subproviders(
|
|||||||
try:
|
try:
|
||||||
payload = resp.json()
|
payload = resp.json()
|
||||||
data = payload.get("data") if isinstance(payload, dict) else None
|
data = payload.get("data") if isinstance(payload, dict) else None
|
||||||
endpoints = data.get("endpoints") if isinstance(data, dict) else None
|
if not isinstance(data, dict):
|
||||||
|
raise ValueError("data must be an object")
|
||||||
|
endpoints = data.get("endpoints")
|
||||||
if not isinstance(endpoints, list):
|
if not isinstance(endpoints, list):
|
||||||
raise ValueError("endpoints must be a list")
|
raise ValueError("endpoints must be a list")
|
||||||
|
common_metadata = {
|
||||||
|
key: value for key, value in data.items() if key != "endpoints"
|
||||||
|
}
|
||||||
identities: dict[str, EndpointIdentity] = {}
|
identities: dict[str, EndpointIdentity] = {}
|
||||||
for endpoint in endpoints:
|
for endpoint in endpoints:
|
||||||
if not isinstance(endpoint, dict):
|
if not isinstance(endpoint, dict):
|
||||||
@@ -281,6 +333,7 @@ async def _fetch_openrouter_endpoint_subproviders(
|
|||||||
provider_name=provider_name
|
provider_name=provider_name
|
||||||
if isinstance(provider_name, str) and provider_name
|
if isinstance(provider_name, str) and provider_name
|
||||||
else None,
|
else None,
|
||||||
|
model_metadata={**common_metadata, **endpoint},
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
if endpoints and not identities:
|
if endpoints and not identities:
|
||||||
@@ -341,6 +394,34 @@ async def _load_model_visibility() -> tuple[
|
|||||||
return overrides_by_key, disabled_model_keys, provider_identities
|
return overrides_by_key, disabled_model_keys, provider_identities
|
||||||
|
|
||||||
|
|
||||||
|
def _serialize_model_metadata(model: object, model_id: str) -> dict[str, Any]:
|
||||||
|
"""Serialize provider-specific model details into the public API shape."""
|
||||||
|
model_dict = getattr(model, "dict", None)
|
||||||
|
if callable(model_dict):
|
||||||
|
metadata = dict(model_dict())
|
||||||
|
else:
|
||||||
|
metadata = {
|
||||||
|
key: value for key, value in vars(model).items() if not key.startswith("_")
|
||||||
|
}
|
||||||
|
|
||||||
|
for field in (
|
||||||
|
"architecture",
|
||||||
|
"pricing",
|
||||||
|
"sats_pricing",
|
||||||
|
"per_request_limits",
|
||||||
|
"top_provider",
|
||||||
|
"alias_ids",
|
||||||
|
):
|
||||||
|
value = metadata.get(field)
|
||||||
|
if isinstance(value, str):
|
||||||
|
try:
|
||||||
|
metadata[field] = json.loads(value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
pass
|
||||||
|
metadata["id"] = model_id
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
|
||||||
def _apply_model_visibility(
|
def _apply_model_visibility(
|
||||||
upstream: BaseUpstreamProvider,
|
upstream: BaseUpstreamProvider,
|
||||||
overrides_by_key: dict[ModelKey, ModelRow] | None,
|
overrides_by_key: dict[ModelKey, ModelRow] | None,
|
||||||
@@ -348,11 +429,10 @@ def _apply_model_visibility(
|
|||||||
) -> list[object]:
|
) -> list[object]:
|
||||||
"""Return provider models after DB disabled/override state is applied.
|
"""Return provider models after DB disabled/override state is applied.
|
||||||
|
|
||||||
Only the identity fields (``id``, ``forwarded_model_id``,
|
DB override rows are used directly rather than rebuilt into priced
|
||||||
``canonical_slug``) matter for path discovery, so DB override rows are used
|
``Model`` objects. Their JSON metadata fields are decoded when each path is
|
||||||
directly rather than rebuilt into fully priced ``Model`` objects — the
|
collected, preserving the provider-specific stored values without running
|
||||||
pricing pipeline costs ~0.7ms of event-loop CPU per row for data this
|
the routing price-selection pipeline.
|
||||||
module immediately discards.
|
|
||||||
"""
|
"""
|
||||||
overrides_by_key = overrides_by_key or {}
|
overrides_by_key = overrides_by_key or {}
|
||||||
disabled_model_keys = disabled_model_keys or set()
|
disabled_model_keys = disabled_model_keys or set()
|
||||||
@@ -414,6 +494,7 @@ async def _collect_provider_paths(
|
|||||||
provider_identity.base_url, provider_identity.id, model_id
|
provider_identity.base_url, provider_identity.id, model_id
|
||||||
),
|
),
|
||||||
provider=provider_identity,
|
provider=provider_identity,
|
||||||
|
model_metadata=_serialize_model_metadata(model, model_id),
|
||||||
)
|
)
|
||||||
|
|
||||||
if not is_openrouter_base_url(upstream.base_url):
|
if not is_openrouter_base_url(upstream.base_url):
|
||||||
@@ -453,6 +534,7 @@ async def _collect_provider_paths(
|
|||||||
endpoint.tag,
|
endpoint.tag,
|
||||||
),
|
),
|
||||||
provider=provider_identity,
|
provider=provider_identity,
|
||||||
|
model_metadata={**endpoint.model_metadata, "id": model_id},
|
||||||
endpoint_tag=endpoint.tag,
|
endpoint_tag=endpoint.tag,
|
||||||
endpoint_name=endpoint.provider_name,
|
endpoint_name=endpoint.provider_name,
|
||||||
)
|
)
|
||||||
@@ -512,6 +594,7 @@ async def _persist_provider_paths(
|
|||||||
"provider_type": discovered.provider.provider_type,
|
"provider_type": discovered.provider.provider_type,
|
||||||
"endpoint_tag": discovered.endpoint_tag,
|
"endpoint_tag": discovered.endpoint_tag,
|
||||||
"endpoint_name": discovered.endpoint_name,
|
"endpoint_name": discovered.endpoint_name,
|
||||||
|
"model_metadata": json.dumps(discovered.model_metadata),
|
||||||
"upstream_provider_id": upstream_provider_id,
|
"upstream_provider_id": upstream_provider_id,
|
||||||
"updated_at": now,
|
"updated_at": now,
|
||||||
}
|
}
|
||||||
@@ -526,6 +609,7 @@ async def _persist_provider_paths(
|
|||||||
"provider_type": insert_stmt.excluded.provider_type,
|
"provider_type": insert_stmt.excluded.provider_type,
|
||||||
"endpoint_tag": insert_stmt.excluded.endpoint_tag,
|
"endpoint_tag": insert_stmt.excluded.endpoint_tag,
|
||||||
"endpoint_name": insert_stmt.excluded.endpoint_name,
|
"endpoint_name": insert_stmt.excluded.endpoint_name,
|
||||||
|
"model_metadata": insert_stmt.excluded.model_metadata,
|
||||||
"updated_at": insert_stmt.excluded.updated_at,
|
"updated_at": insert_stmt.excluded.updated_at,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
@@ -736,10 +820,83 @@ async def refresh_model_paths_periodically(
|
|||||||
break
|
break
|
||||||
|
|
||||||
|
|
||||||
def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
|
def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None:
|
||||||
|
"""Run a path's USD rates through the ``/v1/models`` pricing pipeline.
|
||||||
|
|
||||||
|
Metadata copied from the provider model cache is already priced. OpenRouter
|
||||||
|
endpoint metadata is not: it carries that endpoint's own USD rates, which
|
||||||
|
still need the cache backfill, the provider fee and the sats conversion.
|
||||||
|
"""
|
||||||
|
pricing = model.get("pricing")
|
||||||
|
if model.get("sats_pricing") or not isinstance(pricing, dict):
|
||||||
|
return
|
||||||
|
|
||||||
|
from ..payment.models import (
|
||||||
|
Architecture,
|
||||||
|
Model,
|
||||||
|
Pricing,
|
||||||
|
TopProvider,
|
||||||
|
_calculate_usd_max_costs,
|
||||||
|
_update_model_sats_pricing,
|
||||||
|
backfill_cache_pricing,
|
||||||
|
)
|
||||||
|
from ..payment.price import sats_usd_price
|
||||||
|
|
||||||
|
try:
|
||||||
|
model_id = model.get("forwarded_model_id") or model["id"]
|
||||||
|
usd = backfill_cache_pricing(model_id, Pricing.parse_obj(pricing))
|
||||||
|
usd = Pricing.parse_obj({k: v * provider_fee for k, v in usd.dict().items()})
|
||||||
|
priced = Model(
|
||||||
|
id=model_id,
|
||||||
|
name=model.get("name") or model_id,
|
||||||
|
created=0,
|
||||||
|
description="",
|
||||||
|
context_length=model.get("context_length") or 0,
|
||||||
|
architecture=Architecture(
|
||||||
|
modality="text",
|
||||||
|
input_modalities=[],
|
||||||
|
output_modalities=[],
|
||||||
|
tokenizer="",
|
||||||
|
instruct_type=None,
|
||||||
|
),
|
||||||
|
pricing=usd,
|
||||||
|
top_provider=TopProvider(
|
||||||
|
context_length=model.get("context_length"),
|
||||||
|
max_completion_tokens=model.get("max_completion_tokens"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
(
|
||||||
|
usd.max_prompt_cost,
|
||||||
|
usd.max_completion_cost,
|
||||||
|
usd.max_cost,
|
||||||
|
) = _calculate_usd_max_costs(priced)
|
||||||
|
priced = _update_model_sats_pricing(priced, sats_usd_price())
|
||||||
|
except Exception as exc:
|
||||||
|
# An endpoint with rates we cannot price is still a usable route, so it
|
||||||
|
# is served with its raw upstream pricing rather than dropped.
|
||||||
|
logger.warning(
|
||||||
|
"Could not calculate sats pricing for model path",
|
||||||
|
extra={"model_id": model.get("id"), "error": str(exc)},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if priced.sats_pricing:
|
||||||
|
model["pricing"] = usd.dict()
|
||||||
|
model["sats_pricing"] = priced.sats_pricing.dict()
|
||||||
|
|
||||||
|
|
||||||
|
def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]:
|
||||||
endpoint = None
|
endpoint = None
|
||||||
if row.endpoint_tag or row.endpoint_name:
|
if row.endpoint_tag or row.endpoint_name:
|
||||||
endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name}
|
endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name}
|
||||||
|
try:
|
||||||
|
model = json.loads(row.model_metadata)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
model = {}
|
||||||
|
if not isinstance(model, dict):
|
||||||
|
model = {}
|
||||||
|
model.setdefault("id", row.model_id)
|
||||||
|
_price_in_sats(model, provider_fee)
|
||||||
return {
|
return {
|
||||||
"path": row.path,
|
"path": row.path,
|
||||||
"provider": {
|
"provider": {
|
||||||
@@ -748,11 +905,17 @@ def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
|
|||||||
"type": row.provider_type,
|
"type": row.provider_type,
|
||||||
},
|
},
|
||||||
"endpoint": endpoint,
|
"endpoint": endpoint,
|
||||||
|
"model": model,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _provider_fees(session: "AsyncSession") -> dict[int, float]:
|
||||||
|
rows = (await session.exec(select(UpstreamProviderRow))).all()
|
||||||
|
return {row.id: row.provider_fee for row in rows if row.id is not None}
|
||||||
|
|
||||||
|
|
||||||
async def get_all_model_paths() -> dict:
|
async def get_all_model_paths() -> dict:
|
||||||
"""All models with their exact selectable routes."""
|
"""All models with exact routes and provider-specific model metadata."""
|
||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
rows = (
|
rows = (
|
||||||
await session.exec(
|
await session.exec(
|
||||||
@@ -763,6 +926,7 @@ async def get_all_model_paths() -> dict:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
).all()
|
).all()
|
||||||
|
fees = await _provider_fees(session)
|
||||||
|
|
||||||
grouped: dict[str, list[dict[str, Any]]] = {}
|
grouped: dict[str, list[dict[str, Any]]] = {}
|
||||||
seen_paths: dict[str, set[str]] = {}
|
seen_paths: dict[str, set[str]] = {}
|
||||||
@@ -772,7 +936,9 @@ async def get_all_model_paths() -> dict:
|
|||||||
if row.path in seen_paths.setdefault(row.model_id, set()):
|
if row.path in seen_paths.setdefault(row.model_id, set()):
|
||||||
continue
|
continue
|
||||||
seen_paths[row.model_id].add(row.path)
|
seen_paths[row.model_id].add(row.path)
|
||||||
grouped.setdefault(row.model_id, []).append(_serialize_path(row))
|
grouped.setdefault(row.model_id, []).append(
|
||||||
|
_serialize_path(row, fees.get(row.upstream_provider_id, 1.01))
|
||||||
|
)
|
||||||
data = [
|
data = [
|
||||||
{
|
{
|
||||||
"id": grouped_model_id,
|
"id": grouped_model_id,
|
||||||
@@ -806,6 +972,7 @@ async def get_paths_for_model(model_id: str) -> dict:
|
|||||||
unprefixed_id = public_model_id(model_id)
|
unprefixed_id = public_model_id(model_id)
|
||||||
if unprefixed_id != model_id:
|
if unprefixed_id != model_id:
|
||||||
rows = await load_rows(session, unprefixed_id)
|
rows = await load_rows(session, unprefixed_id)
|
||||||
|
fees = await _provider_fees(session)
|
||||||
|
|
||||||
seen: set[str] = set()
|
seen: set[str] = set()
|
||||||
paths: list[dict] = []
|
paths: list[dict] = []
|
||||||
@@ -815,5 +982,5 @@ async def get_paths_for_model(model_id: str) -> dict:
|
|||||||
if row.path in seen:
|
if row.path in seen:
|
||||||
continue
|
continue
|
||||||
seen.add(row.path)
|
seen.add(row.path)
|
||||||
paths.append(_serialize_path(row))
|
paths.append(_serialize_path(row, fees.get(row.upstream_provider_id, 1.01)))
|
||||||
return {"data": paths, "updated_at": updated_at or None}
|
return {"data": paths, "updated_at": updated_at or None}
|
||||||
|
|||||||
@@ -66,9 +66,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
|||||||
"""Strip 'ollama/' prefix for Ollama API compatibility."""
|
"""Strip 'ollama/' prefix for Ollama API compatibility."""
|
||||||
return model_id.removeprefix("ollama/")
|
return model_id.removeprefix("ollama/")
|
||||||
|
|
||||||
def get_request_base_url(
|
def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str:
|
||||||
self, path: str, model_obj: Model | None = None
|
|
||||||
) -> str:
|
|
||||||
"""Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
|
"""Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
|
||||||
return f"{self.base_url.rstrip('/')}/v1"
|
return f"{self.base_url.rstrip('/')}/v1"
|
||||||
|
|
||||||
@@ -185,7 +183,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
|||||||
except Exception:
|
except Exception:
|
||||||
self._models_cache = models_with_fees
|
self._models_cache = models_with_fees
|
||||||
|
|
||||||
self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache}
|
self._models_by_id = {
|
||||||
|
m.forwarded_model_id or m.id: m for m in self._models_cache
|
||||||
|
}
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Refreshed models cache for {self.base_url}",
|
f"Refreshed models cache for {self.base_url}",
|
||||||
extra={"model_count": len(models)},
|
extra={"model_count": len(models)},
|
||||||
@@ -224,26 +224,14 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
|||||||
Returns:
|
Returns:
|
||||||
Model with provider fee applied to pricing and max costs calculated
|
Model with provider fee applied to pricing and max costs calculated
|
||||||
"""
|
"""
|
||||||
from ..payment.models import Model, Pricing, _calculate_usd_max_costs
|
from ..payment.models import Pricing, _calculate_usd_max_costs
|
||||||
|
|
||||||
adjusted_pricing = Pricing.parse_obj(
|
adjusted_pricing = Pricing.parse_obj(
|
||||||
{k: v * self.provider_fee for k, v in model.pricing.dict().items()}
|
{k: v * self.provider_fee for k, v in model.pricing.dict().items()}
|
||||||
)
|
)
|
||||||
|
|
||||||
temp_model = Model(
|
temp_model = model.copy(
|
||||||
id=model.id,
|
update={"pricing": adjusted_pricing, "sats_pricing": None}
|
||||||
name=model.name,
|
|
||||||
created=model.created,
|
|
||||||
description=model.description,
|
|
||||||
context_length=model.context_length,
|
|
||||||
architecture=model.architecture,
|
|
||||||
pricing=adjusted_pricing,
|
|
||||||
sats_pricing=None,
|
|
||||||
per_request_limits=model.per_request_limits,
|
|
||||||
top_provider=model.top_provider,
|
|
||||||
enabled=model.enabled,
|
|
||||||
upstream_provider_id=model.upstream_provider_id,
|
|
||||||
canonical_slug=model.canonical_slug,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
(
|
(
|
||||||
@@ -252,18 +240,4 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
|||||||
adjusted_pricing.max_cost,
|
adjusted_pricing.max_cost,
|
||||||
) = _calculate_usd_max_costs(temp_model)
|
) = _calculate_usd_max_costs(temp_model)
|
||||||
|
|
||||||
return Model(
|
return model.copy(update={"pricing": adjusted_pricing})
|
||||||
id=model.id,
|
|
||||||
name=model.name,
|
|
||||||
created=model.created,
|
|
||||||
description=model.description,
|
|
||||||
context_length=model.context_length,
|
|
||||||
architecture=model.architecture,
|
|
||||||
pricing=adjusted_pricing,
|
|
||||||
sats_pricing=model.sats_pricing,
|
|
||||||
per_request_limits=model.per_request_limits,
|
|
||||||
top_provider=model.top_provider,
|
|
||||||
enabled=model.enabled,
|
|
||||||
upstream_provider_id=model.upstream_provider_id,
|
|
||||||
canonical_slug=model.canonical_slug,
|
|
||||||
)
|
|
||||||
|
|||||||
+92
-20
@@ -1,5 +1,9 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import random
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass, field
|
||||||
from typing import TYPE_CHECKING, Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@@ -15,6 +19,87 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
_PPQ_SAFE_READ_ATTEMPTS = 3
|
||||||
|
_PPQ_CIRCUIT_COOLDOWN_SECONDS = 30.0
|
||||||
|
|
||||||
|
|
||||||
|
class PPQCircuitOpenError(RuntimeError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _PPQCircuitState:
|
||||||
|
consecutive_failures: int = 0
|
||||||
|
cooldown_until: float = 0.0
|
||||||
|
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||||
|
loop: asyncio.AbstractEventLoop | None = None
|
||||||
|
|
||||||
|
|
||||||
|
_ppq_circuits: dict[str, _PPQCircuitState] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def _ppq_origin(url: str) -> str:
|
||||||
|
parsed = httpx.URL(url)
|
||||||
|
port = parsed.port or {"https": 443, "http": 80}.get(parsed.scheme, 0)
|
||||||
|
return f"{parsed.scheme}://{parsed.host}:{port}"
|
||||||
|
|
||||||
|
|
||||||
|
async def _safe_read_request(
|
||||||
|
client: httpx.AsyncClient,
|
||||||
|
method: str,
|
||||||
|
url: str,
|
||||||
|
*,
|
||||||
|
headers: dict[str, str],
|
||||||
|
json: dict[str, object] | None = None,
|
||||||
|
) -> httpx.Response:
|
||||||
|
state = _ppq_circuits.setdefault(_ppq_origin(url), _PPQCircuitState())
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
if state.loop is not loop:
|
||||||
|
# Locks cannot be reused across event loops.
|
||||||
|
state.lock = asyncio.Lock()
|
||||||
|
state.loop = loop
|
||||||
|
async with state.lock:
|
||||||
|
remaining = state.cooldown_until - time.monotonic()
|
||||||
|
if remaining > 0:
|
||||||
|
raise PPQCircuitOpenError(
|
||||||
|
f"PPQ.AI safe-read circuit is open; retry after {remaining:.2f}s"
|
||||||
|
)
|
||||||
|
|
||||||
|
for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1):
|
||||||
|
try:
|
||||||
|
response = await client.request(method, url, headers=headers, json=json)
|
||||||
|
response.raise_for_status()
|
||||||
|
state.consecutive_failures = 0
|
||||||
|
state.cooldown_until = 0.0
|
||||||
|
return response
|
||||||
|
except (httpx.TransportError, httpx.HTTPStatusError) as error:
|
||||||
|
retryable_status = isinstance(error, httpx.HTTPStatusError) and (
|
||||||
|
error.response.status_code in {502, 503, 504}
|
||||||
|
)
|
||||||
|
if not isinstance(error, httpx.TransportError) and not retryable_status:
|
||||||
|
raise
|
||||||
|
state.consecutive_failures += 1
|
||||||
|
if attempt >= _PPQ_SAFE_READ_ATTEMPTS:
|
||||||
|
state.cooldown_until = (
|
||||||
|
time.monotonic() + _PPQ_CIRCUIT_COOLDOWN_SECONDS
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
base_delay = 0.25 * (2 ** (attempt - 1))
|
||||||
|
delay = base_delay + random.uniform(0.0, base_delay)
|
||||||
|
logger.warning(
|
||||||
|
"PPQ.AI safe read failed; retrying",
|
||||||
|
extra={
|
||||||
|
"url": url,
|
||||||
|
"attempt": attempt,
|
||||||
|
"max_attempts": _PPQ_SAFE_READ_ATTEMPTS,
|
||||||
|
"backoff_seconds": round(delay, 3),
|
||||||
|
"error": repr(error),
|
||||||
|
"error_type": type(error).__name__,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
await asyncio.sleep(delay)
|
||||||
|
raise RuntimeError("unreachable")
|
||||||
|
|
||||||
|
|
||||||
class PPQAIModelPricing(BaseModel):
|
class PPQAIModelPricing(BaseModel):
|
||||||
ui: Optional[dict[str, float]] = None
|
ui: Optional[dict[str, float]] = None
|
||||||
@@ -123,10 +208,8 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
|||||||
url = f"{self.base_url}/models"
|
url = f"{self.base_url}/models"
|
||||||
headers = {"Authorization": f"Bearer {self.api_key}"}
|
headers = {"Authorization": f"Bearer {self.api_key}"}
|
||||||
|
|
||||||
try:
|
|
||||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||||
response = await client.get(url, headers=headers)
|
response = await _safe_read_request(client, "GET", url, headers=headers)
|
||||||
response.raise_for_status()
|
|
||||||
data = response.json()
|
data = response.json()
|
||||||
|
|
||||||
models_data = data.get("data", [])
|
models_data = data.get("data", [])
|
||||||
@@ -157,9 +240,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
|||||||
if or_model:
|
if or_model:
|
||||||
input_price = None
|
input_price = None
|
||||||
if ppqai_model.pricing.api:
|
if ppqai_model.pricing.api:
|
||||||
input_price = ppqai_model.pricing.api.get(
|
input_price = ppqai_model.pricing.api.get("input_per_1M")
|
||||||
"input_per_1M"
|
|
||||||
)
|
|
||||||
elif ppqai_model.pricing.input_per_1M_tokens:
|
elif ppqai_model.pricing.input_per_1M_tokens:
|
||||||
input_price = ppqai_model.pricing.input_per_1M_tokens
|
input_price = ppqai_model.pricing.input_per_1M_tokens
|
||||||
|
|
||||||
@@ -168,9 +249,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
|||||||
|
|
||||||
output_price = None
|
output_price = None
|
||||||
if ppqai_model.pricing.api:
|
if ppqai_model.pricing.api:
|
||||||
output_price = ppqai_model.pricing.api.get(
|
output_price = ppqai_model.pricing.api.get("output_per_1M")
|
||||||
"output_per_1M"
|
|
||||||
)
|
|
||||||
elif ppqai_model.pricing.output_per_1M_tokens:
|
elif ppqai_model.pricing.output_per_1M_tokens:
|
||||||
output_price = ppqai_model.pricing.output_per_1M_tokens
|
output_price = ppqai_model.pricing.output_per_1M_tokens
|
||||||
|
|
||||||
@@ -233,13 +312,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
|||||||
|
|
||||||
return models
|
return models
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(
|
|
||||||
"Error fetching models from PPQ.AI",
|
|
||||||
extra={"error": str(e), "error_type": type(e).__name__},
|
|
||||||
)
|
|
||||||
return []
|
|
||||||
|
|
||||||
async def on_upstream_error_redirect(
|
async def on_upstream_error_redirect(
|
||||||
self, status_code: int, error_message: str
|
self, status_code: int, error_message: str
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -360,8 +432,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||||
response = await client.get(url, headers=headers)
|
response = await _safe_read_request(client, "GET", url, headers=headers)
|
||||||
response.raise_for_status()
|
|
||||||
status_data = response.json()
|
status_data = response.json()
|
||||||
|
|
||||||
is_paid = status_data.get("status") == "Settled"
|
is_paid = status_data.get("status") == "Settled"
|
||||||
@@ -460,8 +531,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
|||||||
logger.debug("Checking PPQ.AI account balance", extra={"url": url})
|
logger.debug("Checking PPQ.AI account balance", extra={"url": url})
|
||||||
|
|
||||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||||
response = await client.post(url, headers=headers, json={})
|
response = await _safe_read_request(
|
||||||
response.raise_for_status()
|
client, "POST", url, headers=headers, json={}
|
||||||
|
)
|
||||||
balance_data = response.json()
|
balance_data = response.json()
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
|
|||||||
@@ -17,6 +17,8 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
|
from ..payment.rates import coerce_rate
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ResolvedPricing:
|
class ResolvedPricing:
|
||||||
@@ -65,11 +67,8 @@ def estimate_context_length(model_id: str) -> int:
|
|||||||
|
|
||||||
|
|
||||||
def _as_float(value: object) -> float | None:
|
def _as_float(value: object) -> float | None:
|
||||||
"""OpenRouter reports prices as strings; coerce, ``None`` if unparseable."""
|
"""OpenRouter reports prices as strings; coerce, ``None`` if not a real rate."""
|
||||||
try:
|
return coerce_rate(value)
|
||||||
return float(value) # type: ignore[arg-type]
|
|
||||||
except (TypeError, ValueError):
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _as_int(value: object) -> int | None:
|
def _as_int(value: object) -> int | None:
|
||||||
@@ -78,23 +77,24 @@ def _as_int(value: object) -> int | None:
|
|||||||
|
|
||||||
|
|
||||||
def _from_litellm(model_id: str) -> ResolvedPricing | None:
|
def _from_litellm(model_id: str) -> ResolvedPricing | None:
|
||||||
# Lazy import so the resolver stays import-light and shares the exact
|
# Lazy import so the resolver shares the exact lookup semantics used by
|
||||||
# lookup semantics used by cache-rate backfill.
|
# cache-rate backfill without importing the models module at load time.
|
||||||
from ..payment.models import litellm_cost_entry
|
from ..payment.models import litellm_cost_entry
|
||||||
|
|
||||||
info = litellm_cost_entry(model_id)
|
info = litellm_cost_entry(model_id)
|
||||||
if info is None:
|
if info is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
prompt = info.get("input_cost_per_token")
|
prompt = coerce_rate(info.get("input_cost_per_token"))
|
||||||
completion = info.get("output_cost_per_token")
|
completion = coerce_rate(info.get("output_cost_per_token"))
|
||||||
if not isinstance(prompt, (int, float)) or not isinstance(completion, (int, float)):
|
if prompt is None or completion is None:
|
||||||
return None
|
return None
|
||||||
# A both-zero entry is litellm listing a model without a real price (free
|
# 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
|
# 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
|
# the model for free. Reject it so the caller falls through, mirroring
|
||||||
# through, mirroring async_fetch_openrouter_models' _has_valid_pricing.
|
# async_fetch_openrouter_models' _has_valid_pricing. Coercion runs first:
|
||||||
if prompt < 0 or completion < 0 or (prompt == 0 and completion == 0):
|
# `NaN` would defeat this guard on its own, every comparison being False.
|
||||||
|
if prompt == 0 and completion == 0:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
input_modalities = ["text"]
|
input_modalities = ["text"]
|
||||||
@@ -102,8 +102,8 @@ def _from_litellm(model_id: str) -> ResolvedPricing | None:
|
|||||||
input_modalities.append("image")
|
input_modalities.append("image")
|
||||||
|
|
||||||
return ResolvedPricing(
|
return ResolvedPricing(
|
||||||
prompt=float(prompt),
|
prompt=prompt,
|
||||||
completion=float(completion),
|
completion=completion,
|
||||||
# max_input_tokens is the context window; max_tokens is litellm's
|
# max_input_tokens is the context window; max_tokens is litellm's
|
||||||
# completion cap (it tracks max_output_tokens for ~94% of models), so
|
# 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
|
# it is never a context source. A missing window falls to the id-based
|
||||||
|
|||||||
@@ -0,0 +1,197 @@
|
|||||||
|
"""Map client reasoning/thinking effort onto a model's allowlist.
|
||||||
|
|
||||||
|
OpenRouter (and some other catalogs) publish per-model reasoning metadata:
|
||||||
|
which effort levels are legal, the default, and whether reasoning is
|
||||||
|
mandatory. Clients still send the generic OpenAI / Anthropic shapes
|
||||||
|
(``reasoning_effort``, ``reasoning.effort``, ``thinking``). This module
|
||||||
|
normalizes those into a supported effort and writes the fields the
|
||||||
|
upstream actually accepts, instead of dropping the parameter or
|
||||||
|
forwarding a value the model rejects.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from ..payment.models import Model, Reasoning
|
||||||
|
|
||||||
|
# Highest first. Unknown values are treated as unranked.
|
||||||
|
EFFORT_RANK: tuple[str, ...] = (
|
||||||
|
"max",
|
||||||
|
"xhigh",
|
||||||
|
"high",
|
||||||
|
"medium",
|
||||||
|
"low",
|
||||||
|
"minimal",
|
||||||
|
"none",
|
||||||
|
)
|
||||||
|
_RANK_INDEX: dict[str, int] = {name: i for i, name in enumerate(EFFORT_RANK)}
|
||||||
|
|
||||||
|
_REASONING_KEYS = ("reasoning", "reasoning_effort", "thinking")
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_effort(value: object) -> str | None:
|
||||||
|
if not isinstance(value, str):
|
||||||
|
return None
|
||||||
|
cleaned = value.strip().lower()
|
||||||
|
return cleaned or None
|
||||||
|
|
||||||
|
|
||||||
|
def closest_supported_effort(
|
||||||
|
requested: str | None,
|
||||||
|
supported: list[str],
|
||||||
|
*,
|
||||||
|
default_effort: str | None = None,
|
||||||
|
mandatory: bool = False,
|
||||||
|
) -> str | None:
|
||||||
|
"""Pick a legal effort for ``requested``.
|
||||||
|
|
||||||
|
Exact match wins. Otherwise the nearest rank in ``EFFORT_RANK`` is
|
||||||
|
used (preferring the higher neighbour on a tie). ``none`` is rejected
|
||||||
|
when ``mandatory`` is set. Missing / unmapped requests fall back to
|
||||||
|
``default_effort``, then the highest remaining supported level.
|
||||||
|
"""
|
||||||
|
allowed_efforts: list[str] = [
|
||||||
|
normalized
|
||||||
|
for item in supported
|
||||||
|
if (normalized := _normalize_effort(item)) is not None
|
||||||
|
]
|
||||||
|
if mandatory:
|
||||||
|
allowed_efforts = [item for item in allowed_efforts if item != "none"]
|
||||||
|
if not allowed_efforts:
|
||||||
|
if mandatory:
|
||||||
|
return _normalize_effort(default_effort)
|
||||||
|
return _normalize_effort(requested) or _normalize_effort(default_effort)
|
||||||
|
|
||||||
|
default = _normalize_effort(default_effort)
|
||||||
|
if default not in allowed_efforts:
|
||||||
|
default = allowed_efforts[0]
|
||||||
|
|
||||||
|
requested_norm = _normalize_effort(requested)
|
||||||
|
if requested_norm is None or (requested_norm == "none" and mandatory):
|
||||||
|
return default
|
||||||
|
|
||||||
|
if requested_norm in allowed_efforts:
|
||||||
|
return requested_norm
|
||||||
|
|
||||||
|
if requested_norm not in _RANK_INDEX:
|
||||||
|
return default
|
||||||
|
|
||||||
|
target = _RANK_INDEX[requested_norm]
|
||||||
|
return min(
|
||||||
|
allowed_efforts,
|
||||||
|
key=lambda effort: (
|
||||||
|
abs(_RANK_INDEX.get(effort, 10_000) - target),
|
||||||
|
_RANK_INDEX.get(effort, 10_000),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_effort(requested: str | None, reasoning: Reasoning | None) -> str | None:
|
||||||
|
"""Map ``requested`` through ``reasoning`` metadata when present."""
|
||||||
|
if reasoning is None:
|
||||||
|
return _normalize_effort(requested)
|
||||||
|
supported = reasoning.supported_efforts or []
|
||||||
|
return closest_supported_effort(
|
||||||
|
requested,
|
||||||
|
supported,
|
||||||
|
default_effort=reasoning.default_effort,
|
||||||
|
mandatory=bool(reasoning.mandatory),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _effort_from_thinking(thinking: object) -> str | None:
|
||||||
|
if not isinstance(thinking, dict):
|
||||||
|
return None
|
||||||
|
effort = _normalize_effort(thinking.get("effort"))
|
||||||
|
if effort:
|
||||||
|
return effort
|
||||||
|
thinking_type = _normalize_effort(thinking.get("type"))
|
||||||
|
if thinking_type in {"disabled", "none"}:
|
||||||
|
return "none"
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def extract_requested_effort(data: dict[str, Any]) -> str | None:
|
||||||
|
"""Best-effort effort string from the OpenAI / Anthropic request shapes."""
|
||||||
|
if isinstance(data.get("reasoning"), dict):
|
||||||
|
nested = _normalize_effort(data["reasoning"].get("effort"))
|
||||||
|
if nested:
|
||||||
|
return nested
|
||||||
|
top_level = _normalize_effort(data.get("reasoning_effort"))
|
||||||
|
if top_level:
|
||||||
|
return top_level
|
||||||
|
return _effort_from_thinking(data.get("thinking"))
|
||||||
|
|
||||||
|
|
||||||
|
def _request_mentions_reasoning(data: dict[str, Any]) -> bool:
|
||||||
|
return any(key in data for key in _REASONING_KEYS)
|
||||||
|
|
||||||
|
|
||||||
|
def apply_reasoning_effort(
|
||||||
|
data: dict[str, Any],
|
||||||
|
model: Model,
|
||||||
|
*,
|
||||||
|
drop_thinking: bool = False,
|
||||||
|
) -> bool:
|
||||||
|
"""Rewrite ``data`` in place so effort matches the model allowlist.
|
||||||
|
|
||||||
|
Returns True when ``data`` changed. Leaves the body alone when the
|
||||||
|
caller did not send a reasoning field and the model does not require
|
||||||
|
one. Existing ``reasoning`` object keys (``max_tokens``, ``exclude``,
|
||||||
|
``enabled``) are preserved; only ``effort`` is mapped.
|
||||||
|
|
||||||
|
``drop_thinking`` is for OpenAI-compatible backends that reject the
|
||||||
|
Anthropic ``thinking`` object: the effort is lifted onto
|
||||||
|
``reasoning_effort`` / ``reasoning.effort`` and ``thinking`` is removed.
|
||||||
|
"""
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
return False
|
||||||
|
|
||||||
|
reasoning_meta = getattr(model, "reasoning", None)
|
||||||
|
mentioned = _request_mentions_reasoning(data)
|
||||||
|
if not mentioned and not (reasoning_meta and reasoning_meta.mandatory):
|
||||||
|
return False
|
||||||
|
|
||||||
|
resolved = resolve_effort(extract_requested_effort(data), reasoning_meta)
|
||||||
|
changed = False
|
||||||
|
|
||||||
|
if drop_thinking and "thinking" in data:
|
||||||
|
data.pop("thinking", None)
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
if resolved is None:
|
||||||
|
return changed
|
||||||
|
|
||||||
|
if "reasoning_effort" in data:
|
||||||
|
if data.get("reasoning_effort") != resolved:
|
||||||
|
data["reasoning_effort"] = resolved
|
||||||
|
changed = True
|
||||||
|
elif drop_thinking or (reasoning_meta and reasoning_meta.mandatory):
|
||||||
|
# Invent the OpenAI-shaped field when we stripped Anthropic
|
||||||
|
# ``thinking``, or when the model will reject a request with no
|
||||||
|
# effort at all.
|
||||||
|
if not isinstance(data.get("reasoning"), dict):
|
||||||
|
data["reasoning_effort"] = resolved
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
existing = data.get("reasoning")
|
||||||
|
if isinstance(existing, dict):
|
||||||
|
if existing.get("effort") != resolved:
|
||||||
|
data["reasoning"] = {**existing, "effort": resolved}
|
||||||
|
changed = True
|
||||||
|
elif reasoning_meta and reasoning_meta.mandatory and "reasoning_effort" not in data:
|
||||||
|
data["reasoning"] = {"effort": resolved}
|
||||||
|
changed = True
|
||||||
|
|
||||||
|
return changed
|
||||||
|
|
||||||
|
|
||||||
|
def adapt_messages_body_for_litellm(data: dict[str, Any], model: Model) -> None:
|
||||||
|
"""Convert Anthropic ``thinking`` into OpenAI-shaped effort for litellm.
|
||||||
|
|
||||||
|
Litellm's Anthropic-messages adapter talking to an OpenAI-compatible
|
||||||
|
upstream will 400 on ``thinking``. Lift the effort onto
|
||||||
|
``reasoning_effort`` and drop the Anthropic-only object.
|
||||||
|
"""
|
||||||
|
apply_reasoning_effort(data, model, drop_thinking=True)
|
||||||
@@ -84,13 +84,33 @@ def extract_error_message(response: Response) -> str:
|
|||||||
return ""
|
return ""
|
||||||
|
|
||||||
|
|
||||||
def strip_unsupported_param(
|
# Spend-shaping fields bound how much work — and therefore cost — the upstream
|
||||||
body: dict, error_message: str
|
# may perform. The reservation was priced with these caps in place; dropping one
|
||||||
) -> tuple[dict, str] | None:
|
# and retrying would let the request run uncapped (or fan out) and bill above the
|
||||||
|
# caller's authorization. When the upstream names one of these, decline the strip
|
||||||
|
# and let the error propagate. Matched case-insensitively.
|
||||||
|
_SPEND_SHAPING_PARAMS = frozenset(
|
||||||
|
{
|
||||||
|
"max_tokens",
|
||||||
|
"max_completion_tokens",
|
||||||
|
"max_output_tokens",
|
||||||
|
"max_tokens_to_sample",
|
||||||
|
"n",
|
||||||
|
"best_of",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None:
|
||||||
"""Drop a top-level param the upstream named as unsupported/deprecated.
|
"""Drop a top-level param the upstream named as unsupported/deprecated.
|
||||||
|
|
||||||
Returns ``(new_body, param)`` (a new dict, original untouched) when the
|
Returns ``(new_body, param)`` (a new dict, original untouched) when the
|
||||||
error names a top-level param present in the body, otherwise ``None``.
|
error names a top-level param present in the body, otherwise ``None``.
|
||||||
|
|
||||||
|
Spend-shaping fields (output caps, fan-out counts) are never stripped:
|
||||||
|
removing one after the reservation was priced would uncap the retry and
|
||||||
|
overcharge. When the upstream names such a field, decline so the original
|
||||||
|
error propagates rather than silently resizing the request's cost.
|
||||||
"""
|
"""
|
||||||
match = _UNSUPPORTED_PARAM_RE.search(error_message)
|
match = _UNSUPPORTED_PARAM_RE.search(error_message)
|
||||||
if not match:
|
if not match:
|
||||||
@@ -98,6 +118,13 @@ def strip_unsupported_param(
|
|||||||
param = match.group("param")
|
param = match.group("param")
|
||||||
if param not in body:
|
if param not in body:
|
||||||
return None
|
return None
|
||||||
|
if param.lower() in _SPEND_SHAPING_PARAMS:
|
||||||
|
logger.warning(
|
||||||
|
"Upstream rejected spend-shaping param '%s'; refusing to strip it "
|
||||||
|
"(retrying uncapped would overcharge) — surfacing the error",
|
||||||
|
param,
|
||||||
|
)
|
||||||
|
return None
|
||||||
new_body = {k: v for k, v in body.items() if k != param}
|
new_body = {k: v for k, v in body.items() if k != param}
|
||||||
return new_body, param
|
return new_body, param
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
@@ -28,6 +28,7 @@ logger = get_logger(__name__)
|
|||||||
class TinfoilModelPricing(BaseModel):
|
class TinfoilModelPricing(BaseModel):
|
||||||
inputTokenPricePer1M: float = 0.0
|
inputTokenPricePer1M: float = 0.0
|
||||||
outputTokenPricePer1M: float = 0.0
|
outputTokenPricePer1M: float = 0.0
|
||||||
|
cachedInputTokenPricePer1M: Optional[float] = None
|
||||||
requestPrice: float = 0.0
|
requestPrice: float = 0.0
|
||||||
|
|
||||||
|
|
||||||
@@ -186,6 +187,14 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider):
|
|||||||
output_price = tf.pricing.outputTokenPricePer1M
|
output_price = tf.pricing.outputTokenPricePer1M
|
||||||
request_price = tf.pricing.requestPrice
|
request_price = tf.pricing.requestPrice
|
||||||
|
|
||||||
|
# Tinfoil bills cache reads at the cached rate when the
|
||||||
|
# model exposes one, otherwise at the full input rate.
|
||||||
|
# Cache writes are never priced separately — a miss is
|
||||||
|
# just regular input prefill.
|
||||||
|
cached_price = tf.pricing.cachedInputTokenPricePer1M
|
||||||
|
if cached_price is None or cached_price <= 0.0:
|
||||||
|
cached_price = input_price
|
||||||
|
|
||||||
modality = "text->text"
|
modality = "text->text"
|
||||||
input_modalities = ["text"]
|
input_modalities = ["text"]
|
||||||
output_modalities = ["text"]
|
output_modalities = ["text"]
|
||||||
@@ -214,6 +223,8 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider):
|
|||||||
image=0.0,
|
image=0.0,
|
||||||
web_search=0.0,
|
web_search=0.0,
|
||||||
internal_reasoning=0.0,
|
internal_reasoning=0.0,
|
||||||
|
input_cache_read=cached_price / 1_000_000,
|
||||||
|
input_cache_write=input_price / 1_000_000,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -20,11 +20,12 @@ from urllib.parse import urlsplit
|
|||||||
import h11
|
import h11
|
||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
|
from ..core.exceptions import EhbpTimeoutError
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
_READ_BUFSIZE = 65536
|
_READ_BUFSIZE = 65536
|
||||||
_DEFAULT_TIMEOUT_SECONDS = 30.0
|
_DEFAULT_TIMEOUT_SECONDS = 60.0
|
||||||
_DEFAULT_CLOSE_TIMEOUT_SECONDS = 1.0
|
_DEFAULT_CLOSE_TIMEOUT_SECONDS = 1.0
|
||||||
_DEFAULT_MAX_RESPONSE_BYTES = 25 * 1024 * 1024
|
_DEFAULT_MAX_RESPONSE_BYTES = 25 * 1024 * 1024
|
||||||
_HOP_BY_HOP_HEADERS = {
|
_HOP_BY_HOP_HEADERS = {
|
||||||
@@ -100,10 +101,15 @@ async def forward_with_trailer(
|
|||||||
headers = _strip_hop_by_hop_headers(headers)
|
headers = _strip_hop_by_hop_headers(headers)
|
||||||
|
|
||||||
ssl_ctx = ssl.create_default_context()
|
ssl_ctx = ssl.create_default_context()
|
||||||
|
try:
|
||||||
reader, writer = await asyncio.wait_for(
|
reader, writer = await asyncio.wait_for(
|
||||||
asyncio.open_connection(host, port, ssl=ssl_ctx),
|
asyncio.open_connection(host, port, ssl=ssl_ctx),
|
||||||
timeout=timeout_seconds,
|
timeout=timeout_seconds,
|
||||||
)
|
)
|
||||||
|
except asyncio.TimeoutError as exc:
|
||||||
|
raise EhbpTimeoutError(
|
||||||
|
f"EHBP upstream {host} timed out after {timeout_seconds:g}s connecting"
|
||||||
|
) from exc
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Build HTTP/1.1 request
|
# Build HTTP/1.1 request
|
||||||
@@ -126,7 +132,13 @@ async def forward_with_trailer(
|
|||||||
request_data += body
|
request_data += body
|
||||||
|
|
||||||
writer.write(request_data)
|
writer.write(request_data)
|
||||||
|
try:
|
||||||
await asyncio.wait_for(writer.drain(), timeout=timeout_seconds)
|
await asyncio.wait_for(writer.drain(), timeout=timeout_seconds)
|
||||||
|
except asyncio.TimeoutError as exc:
|
||||||
|
raise EhbpTimeoutError(
|
||||||
|
f"EHBP upstream {host} timed out after "
|
||||||
|
f"{timeout_seconds:g}s sending request"
|
||||||
|
) from exc
|
||||||
|
|
||||||
# Parse response with h11
|
# Parse response with h11
|
||||||
conn = h11.Connection(h11.CLIENT)
|
conn = h11.Connection(h11.CLIENT)
|
||||||
@@ -140,10 +152,16 @@ async def forward_with_trailer(
|
|||||||
event = conn.next_event()
|
event = conn.next_event()
|
||||||
|
|
||||||
if event is h11.NEED_DATA:
|
if event is h11.NEED_DATA:
|
||||||
|
try:
|
||||||
data = await asyncio.wait_for(
|
data = await asyncio.wait_for(
|
||||||
reader.read(_READ_BUFSIZE),
|
reader.read(_READ_BUFSIZE),
|
||||||
timeout=timeout_seconds,
|
timeout=timeout_seconds,
|
||||||
)
|
)
|
||||||
|
except asyncio.TimeoutError as exc:
|
||||||
|
raise EhbpTimeoutError(
|
||||||
|
f"EHBP upstream {host} timed out after "
|
||||||
|
f"{timeout_seconds:g}s waiting for response data"
|
||||||
|
) from exc
|
||||||
conn.receive_data(data if data else b"")
|
conn.receive_data(data if data else b"")
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
|||||||
+246
-841
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,89 @@
|
|||||||
|
"""Redeem a cashu token into a balance and pay it out to a Lightning address.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python scripts/refund_token_to_lightning.py <cashu-token> <lightning-address> [--url http://localhost:8000]
|
||||||
|
|
||||||
|
Steps:
|
||||||
|
1. POST /v1/balance/create redeems the token into a fresh API key
|
||||||
|
2. POST /v1/balance/refund pays the full balance to the Lightning address
|
||||||
|
|
||||||
|
A 502 from the refund means the melt was dispatched but unconfirmed; the
|
||||||
|
balance is withheld until the server reconciles it. Re-run with the printed
|
||||||
|
API key to check whether it settled.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import ipaddress
|
||||||
|
import sys
|
||||||
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
|
||||||
|
def _is_loopback(host: str) -> bool:
|
||||||
|
if host == "localhost":
|
||||||
|
return True
|
||||||
|
try:
|
||||||
|
return ipaddress.ip_address(host.strip("[]")).is_loopback
|
||||||
|
except ValueError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def check_url(url: str) -> str:
|
||||||
|
"""Reject a URL that would put the token and the API key on the wire."""
|
||||||
|
parsed = urlparse(url)
|
||||||
|
if parsed.scheme == "https":
|
||||||
|
return url
|
||||||
|
if parsed.scheme == "http" and _is_loopback(parsed.hostname or ""):
|
||||||
|
return url
|
||||||
|
raise SystemExit(
|
||||||
|
f"Refusing to send a cashu token and bearer key to {url!r}: "
|
||||||
|
"use https, or http only for a loopback host."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def create_balance(client: httpx.Client, token: str) -> str:
|
||||||
|
response = client.post("/v1/balance/create", json={"initial_balance_token": token})
|
||||||
|
response.raise_for_status()
|
||||||
|
data = response.json()
|
||||||
|
print(f"Redeemed token: balance {data['balance']} msats, key {data['api_key']}")
|
||||||
|
return str(data["api_key"])
|
||||||
|
|
||||||
|
|
||||||
|
def refund_to_lightning(client: httpx.Client, api_key: str, address: str) -> dict:
|
||||||
|
response = client.post(
|
||||||
|
"/v1/balance/refund",
|
||||||
|
headers={"Authorization": f"Bearer {api_key}"},
|
||||||
|
json={"lightning_address": address},
|
||||||
|
)
|
||||||
|
if response.status_code >= 400:
|
||||||
|
print(f"Refund failed ({response.status_code}): {response.text}")
|
||||||
|
sys.exit(1)
|
||||||
|
return dict(response.json())
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
|
||||||
|
parser.add_argument("token", help="cashu token, or sk-... key from a prior run")
|
||||||
|
parser.add_argument("lightning_address", help="Lightning address or LNURL")
|
||||||
|
parser.add_argument("--url", default="http://localhost:8000", help="routstr URL")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
with httpx.Client(base_url=check_url(args.url), timeout=120.0) as client:
|
||||||
|
api_key = (
|
||||||
|
args.token
|
||||||
|
if args.token.startswith("sk-")
|
||||||
|
else create_balance(client, args.token)
|
||||||
|
)
|
||||||
|
result = refund_to_lightning(client, api_key, args.lightning_address)
|
||||||
|
|
||||||
|
amount = result.get("sats") or result.get("msats")
|
||||||
|
unit = "sats" if "sats" in result else "msats"
|
||||||
|
print(
|
||||||
|
f"Refund {result['refund_id']} {result['status']}: "
|
||||||
|
f"{amount} {unit} -> {result.get('recipient', args.lightning_address)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -7,9 +7,27 @@ absent one) override this per-test via ``monkeypatch``.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
from typing import Iterator
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
# Valid Fernet keys; KEY_A is the suite default, KEY_B is for wrong-key tests.
|
# 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 = "l_Tkp-7xmjcQ-IFhr6qhILrU8HPRbEmYMrfSbo_5srU="
|
||||||
TEST_SECRET_KEY_ALT = "_Teyrky_iToeDK51Tj1FsI9MJ340_cqKGmeher-a7MQ="
|
TEST_SECRET_KEY_ALT = "_Teyrky_iToeDK51Tj1FsI9MJ340_cqKGmeher-a7MQ="
|
||||||
|
|
||||||
os.environ.setdefault("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY)
|
os.environ.setdefault("ROUTSTR_SECRET_KEY", TEST_SECRET_KEY)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _isolate_redemption_negative_cache() -> Iterator[None]:
|
||||||
|
"""Clear the process-wide negative cache between tests.
|
||||||
|
|
||||||
|
The cache deliberately persists terminal redemption failures across
|
||||||
|
requests; without this fixture a test that burns a token would poison
|
||||||
|
every later test reusing the same token string.
|
||||||
|
"""
|
||||||
|
from routstr.redemption_cache import redemption_negative_cache
|
||||||
|
|
||||||
|
redemption_negative_cache.clear()
|
||||||
|
yield
|
||||||
|
redemption_negative_cache.clear()
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
from typing import Any, AsyncGenerator, Callable, Dict, List, Optional, Tuple
|
from typing import Any, AsyncGenerator, Callable, Dict, Iterator, List, Optional, Tuple
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -68,6 +68,14 @@ os.environ.pop("ADMIN_PASSWORD", None)
|
|||||||
|
|
||||||
from routstr.core.db import ApiKey, get_session # noqa: E402
|
from routstr.core.db import ApiKey, get_session # noqa: E402
|
||||||
from routstr.core.main import app, lifespan # noqa: E402
|
from routstr.core.main import app, lifespan # noqa: E402
|
||||||
|
from routstr.mint import MintRateGuard # noqa: E402
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def isolate_mint_rate_guards() -> Iterator[None]:
|
||||||
|
MintRateGuard._guards.clear()
|
||||||
|
yield
|
||||||
|
MintRateGuard._guards.clear()
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session")
|
@pytest.fixture(scope="session")
|
||||||
@@ -508,6 +516,10 @@ async def integration_app(
|
|||||||
# Copy all routes from the main app
|
# Copy all routes from the main app
|
||||||
test_app.router = app.router
|
test_app.router = app.router
|
||||||
|
|
||||||
|
# ...and its exception handlers, so a request that fails here fails the way
|
||||||
|
# it would in production rather than escaping as a bare exception.
|
||||||
|
test_app.exception_handlers.update(app.exception_handlers)
|
||||||
|
|
||||||
# Override the get_session dependency
|
# Override the get_session dependency
|
||||||
async def override_get_session() -> AsyncGenerator[AsyncSession, None]:
|
async def override_get_session() -> AsyncGenerator[AsyncSession, None]:
|
||||||
yield integration_session
|
yield integration_session
|
||||||
@@ -542,8 +554,8 @@ async def integration_app(
|
|||||||
patch("routstr.wallet.send_to_lnurl", testmint_wallet.send_to_lnurl),
|
patch("routstr.wallet.send_to_lnurl", testmint_wallet.send_to_lnurl),
|
||||||
patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token),
|
patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token),
|
||||||
patch("routstr.wallet.get_balance", testmint_wallet.get_balance),
|
patch("routstr.wallet.get_balance", testmint_wallet.get_balance),
|
||||||
patch("routstr.balance.send_token", testmint_wallet.send_token),
|
patch("routstr.refund.send_token", testmint_wallet.send_token),
|
||||||
patch("routstr.balance.send_to_lnurl", testmint_wallet.send_to_lnurl),
|
patch("routstr.refund.send_to_lnurl", testmint_wallet.send_to_lnurl),
|
||||||
patch("websockets.connect") as mock_websockets,
|
patch("websockets.connect") as mock_websockets,
|
||||||
patch("routstr.payment.price.btc_usd_price", return_value=50000.0),
|
patch("routstr.payment.price.btc_usd_price", return_value=50000.0),
|
||||||
patch("routstr.payment.price.sats_usd_price", return_value=0.0005),
|
patch("routstr.payment.price.sats_usd_price", return_value=0.0005),
|
||||||
|
|||||||
@@ -0,0 +1,476 @@
|
|||||||
|
"""Admin write edge: a rate that is not a number never becomes a stored price.
|
||||||
|
|
||||||
|
A billable rate is usable only when it is finite and non-negative. The admin
|
||||||
|
model endpoints are an entry point for rates the node will later bill on, and
|
||||||
|
they accept whatever a client sends: ``json`` parses the bare ``NaN``/
|
||||||
|
``Infinity`` literals into real floats and overflows ``1e999`` to ``inf``, a
|
||||||
|
non-numeric string coerced silently to ``$0``, and a negative rate is truthy so
|
||||||
|
it read back as a chargeable price that bills a negative amount.
|
||||||
|
|
||||||
|
These tests assert the edge answers a malformed rate with a 422 — a client bug
|
||||||
|
reported as a client bug — rather than persisting it or failing as a 500, and
|
||||||
|
that the operator can still open the listing that shows the row needing repair.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from httpx import AsyncClient
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from routstr.core.admin import admin_sessions
|
||||||
|
from routstr.core.db import ModelRow, UpstreamProviderRow
|
||||||
|
from routstr.proxy import reinitialize_upstreams
|
||||||
|
|
||||||
|
|
||||||
|
def _admin_headers() -> dict[str, str]:
|
||||||
|
token = "test-admin-rate-validation-token"
|
||||||
|
admin_sessions[token] = int(
|
||||||
|
(datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp()
|
||||||
|
)
|
||||||
|
return {"Authorization": f"Bearer {token}"}
|
||||||
|
|
||||||
|
|
||||||
|
def _pricing(**overrides: object) -> dict[str, object]:
|
||||||
|
pricing: dict[str, object] = {
|
||||||
|
"prompt": 1.4e-7,
|
||||||
|
"completion": 2.8e-7,
|
||||||
|
"request": 0.0,
|
||||||
|
"image": 0.0,
|
||||||
|
"web_search": 0.0,
|
||||||
|
"internal_reasoning": 0.0,
|
||||||
|
"input_cache_read": 0.0,
|
||||||
|
"input_cache_write": 0.0,
|
||||||
|
}
|
||||||
|
pricing.update(overrides)
|
||||||
|
return pricing
|
||||||
|
|
||||||
|
|
||||||
|
def _payload(
|
||||||
|
provider_id: int,
|
||||||
|
*,
|
||||||
|
model_id: str = "rate-model",
|
||||||
|
pricing: dict[str, object] | None = None,
|
||||||
|
) -> dict[str, object]:
|
||||||
|
return {
|
||||||
|
"id": model_id,
|
||||||
|
"name": "Rate Model",
|
||||||
|
"description": "d",
|
||||||
|
"created": 0,
|
||||||
|
"context_length": 128000,
|
||||||
|
"architecture": {
|
||||||
|
"modality": "text",
|
||||||
|
"input_modalities": ["text"],
|
||||||
|
"output_modalities": ["text"],
|
||||||
|
"tokenizer": "unknown",
|
||||||
|
"instruct_type": None,
|
||||||
|
},
|
||||||
|
"pricing": pricing if pricing is not None else _pricing(),
|
||||||
|
"per_request_limits": None,
|
||||||
|
"top_provider": None,
|
||||||
|
"upstream_provider_id": provider_id,
|
||||||
|
"canonical_slug": None,
|
||||||
|
"alias_ids": [],
|
||||||
|
"enabled": True,
|
||||||
|
"forwarded_model_id": model_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _make_provider(session: AsyncSession) -> int:
|
||||||
|
provider = UpstreamProviderRow(
|
||||||
|
provider_type="generic",
|
||||||
|
base_url="https://rate-upstream.example/v1",
|
||||||
|
api_key="test-key",
|
||||||
|
provider_fee=1.0,
|
||||||
|
)
|
||||||
|
session.add(provider)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(provider)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
assert provider.id is not None
|
||||||
|
return provider.id
|
||||||
|
|
||||||
|
|
||||||
|
def _raw_model_body(provider_id: int, model_id: str, prompt_literal: str) -> str:
|
||||||
|
"""A request body built as text, so it can carry a literal ``json`` accepts
|
||||||
|
but Python's own encoder would refuse to produce."""
|
||||||
|
return (
|
||||||
|
f'{{"id": "{model_id}", "name": "raw", "description": "d", "created": 0,'
|
||||||
|
' "context_length": 8192, "architecture": {"modality": "text"},'
|
||||||
|
f' "pricing": {{"prompt": {prompt_literal}, "completion": 2.8e-7}},'
|
||||||
|
f' "upstream_provider_id": {provider_id}}}'
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_negative_price_is_rejected(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""A negative rate is not a valid price — accepting it would persist a row
|
||||||
|
that bills a negative amount, which settlement subtracts from the balance.
|
||||||
|
Being truthy, it also reads back as a chargeable price. Reject at the edge
|
||||||
|
rather than silently storing it."""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
json=_payload(provider_id, model_id="neg-price", pricing=_pricing(prompt=-1.0)),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 422
|
||||||
|
assert await integration_session.get(ModelRow, ("neg-price", provider_id)) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_malformed_price_string_is_rejected(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""A present non-numeric rate is a client bug: it coerces to ``$0`` on the
|
||||||
|
read path, producing an unpriced-looking row indistinguishable from a
|
||||||
|
deliberate free price. Surface it as a 422 instead of accepting it."""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
json=_payload(
|
||||||
|
provider_id, model_id="bad-price", pricing=_pricing(prompt="oops")
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 422
|
||||||
|
assert await integration_session.get(ModelRow, ("bad-price", provider_id)) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_boolean_price_is_rejected(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""A JSON ``true`` coerces to a finite, positive ``1.0`` — a dollar per
|
||||||
|
token — so it passes every numeric guard. The write edge asks the same
|
||||||
|
coercion the catalog readers do, and answers a 422."""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
json=_payload(
|
||||||
|
provider_id, model_id="bool-price", pricing=_pricing(prompt=True)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 422
|
||||||
|
assert await integration_session.get(ModelRow, ("bool-price", provider_id)) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("model_id", "pricing"),
|
||||||
|
[
|
||||||
|
("null-prompt", _pricing(prompt=None)),
|
||||||
|
("null-aux-rate", _pricing(image=None)),
|
||||||
|
("no-prompt", {k: v for k, v in _pricing().items() if k != "prompt"}),
|
||||||
|
],
|
||||||
|
ids=["null-required", "null-auxiliary", "absent-required"],
|
||||||
|
)
|
||||||
|
async def test_a_rate_that_is_not_there_is_rejected(
|
||||||
|
model_id: str,
|
||||||
|
pricing: dict[str, object],
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A rate given as ``null``, or a required rate left out, is not a price.
|
||||||
|
|
||||||
|
``dict.get`` cannot tell the two apart and skipped both, so a row
|
||||||
|
``Pricing`` cannot parse was committed and the response that reads it back
|
||||||
|
raised.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
json=_payload(provider_id, model_id=model_id, pricing=pricing),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 422
|
||||||
|
assert await integration_session.get(ModelRow, (model_id, provider_id)) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_an_absent_auxiliary_rate_is_still_accepted(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""Only ``prompt`` and ``completion`` are required; the rest carry defaults,
|
||||||
|
and a payload that omits them must still be accepted."""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
json=_payload(
|
||||||
|
provider_id,
|
||||||
|
model_id="lean-price",
|
||||||
|
pricing={"prompt": 1.4e-7, "completion": 2.8e-7},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 200
|
||||||
|
assert await integration_session.get(ModelRow, ("lean-price", provider_id))
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_numeric_string_price_is_still_accepted(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""The stored pricing JSON has always accepted numeric strings, and the UI
|
||||||
|
round-trips rates through text fields. Rejecting a *malformed* rate must not
|
||||||
|
also reject a well-formed one that arrives spelled as a string."""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
json=_payload(
|
||||||
|
provider_id, model_id="string-price", pricing=_pricing(prompt="0.000005")
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 200
|
||||||
|
row = await integration_session.get(ModelRow, ("string-price", provider_id))
|
||||||
|
assert row is not None
|
||||||
|
assert json.loads(row.pricing)["prompt"] == "0.000005"
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_finite_price_is_rejected_by_the_write_model() -> None:
|
||||||
|
"""``NaN``/``±inf`` are not billable rates: the carrier every write endpoint
|
||||||
|
shares must reject them before they can be persisted and read back as a
|
||||||
|
chargeable price."""
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
from routstr.core.admin import ModelCreate
|
||||||
|
|
||||||
|
for bad in (float("nan"), float("inf"), float("-inf")):
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
ModelCreate.model_validate(
|
||||||
|
_payload(1, model_id="nonfinite", pricing=_pricing(prompt=bad))
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_oversized_integer_price_is_rejected(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""A JSON integer too large for a float is a client bug, not a server fault.
|
||||||
|
|
||||||
|
``float()`` raises ``OverflowError`` for it, and pydantic converts only
|
||||||
|
``ValueError``/``AssertionError`` into validation errors, so it escaped the
|
||||||
|
edge as a 500. It must be answered with the same 422 as every other
|
||||||
|
unusable rate.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers={**_admin_headers(), "Content-Type": "application/json"},
|
||||||
|
content=_raw_model_body(provider_id, "huge-price", "9" * 400),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 422
|
||||||
|
assert await integration_session.get(ModelRow, ("huge-price", provider_id)) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_non_finite_literal_price_is_rejected(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""A bare ``Infinity``/``NaN`` literal gets the same 422 as any other rate.
|
||||||
|
|
||||||
|
``json`` accepts both literals, so the edge sees a real float and rejects
|
||||||
|
it — but pydantic echoes the offending value back in the error's ``input``
|
||||||
|
field, and the response encoder runs with ``allow_nan=False``. Serializing
|
||||||
|
that reply raised "Out of range float values are not JSON compliant", so the
|
||||||
|
422 escaped as a 500 and reported a client's bad rate as a server fault.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
for literal in ("Infinity", "-Infinity", "NaN"):
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers={**_admin_headers(), "Content-Type": "application/json"},
|
||||||
|
content=_raw_model_body(provider_id, "odd-price", literal),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 422, literal
|
||||||
|
assert (
|
||||||
|
await integration_session.get(ModelRow, ("odd-price", provider_id)) is None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_non_finite_literal_price_is_rejected_in_batch_override(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""The batch path shares the same carrier, so it must answer 422 too."""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/batch-override",
|
||||||
|
headers={**_admin_headers(), "Content-Type": "application/json"},
|
||||||
|
content=(
|
||||||
|
'{"models": ['
|
||||||
|
+ _raw_model_body(provider_id, "odd-batch", "Infinity")
|
||||||
|
+ "]}"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 422
|
||||||
|
assert await integration_session.get(ModelRow, ("odd-batch", provider_id)) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_admin_model_listing_shows_a_non_finite_stored_rate(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""The operator must be able to see the rate that needs fixing.
|
||||||
|
|
||||||
|
The admin listing is the one view that still carries a row the served
|
||||||
|
catalog holds back, and its encoder rendered a stored ``Infinity`` as
|
||||||
|
``null`` — indistinguishable from a rate the row never carried.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
integration_session.add(
|
||||||
|
ModelRow(
|
||||||
|
id="inf-rate",
|
||||||
|
name="inf-rate",
|
||||||
|
description="d",
|
||||||
|
created=0,
|
||||||
|
context_length=8192,
|
||||||
|
architecture=json.dumps(
|
||||||
|
{
|
||||||
|
"modality": "text",
|
||||||
|
"input_modalities": ["text"],
|
||||||
|
"output_modalities": ["text"],
|
||||||
|
"tokenizer": "unknown",
|
||||||
|
"instruct_type": None,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
pricing=json.dumps({"prompt": float("inf"), "completion": 2e-06}),
|
||||||
|
upstream_provider_id=provider_id,
|
||||||
|
enabled=True,
|
||||||
|
forwarded_model_id="inf-rate",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
resp = await integration_client.get(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 200
|
||||||
|
listed = {m["id"]: m for m in resp.json()["db_models"]}
|
||||||
|
assert listed["inf-rate"]["pricing"]["prompt"] == "inf"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_malformed_auxiliary_rate_is_rejected(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""Validation spans every billable rate, not just the token rates.
|
||||||
|
|
||||||
|
``prompt``/``completion`` are the rates most prices are built from, but the
|
||||||
|
request, image, search, reasoning and cache rates are billed too. A negative
|
||||||
|
or non-finite value in any of them is the same defect and must be answered
|
||||||
|
the same way.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
|
||||||
|
for field, bad in (
|
||||||
|
("request", -1.0),
|
||||||
|
("image", -0.5),
|
||||||
|
("web_search", float("inf")),
|
||||||
|
("internal_reasoning", float("nan")),
|
||||||
|
("input_cache_read", -1e-06),
|
||||||
|
("input_cache_write", float("-inf")),
|
||||||
|
("completion", -1.0),
|
||||||
|
):
|
||||||
|
# Send raw bytes rather than `json=`: httpx>=0.28 refuses to encode
|
||||||
|
# non-finite floats itself (allow_nan=False), but the point of this
|
||||||
|
# test is that the SERVER answers the bare NaN/Infinity literals
|
||||||
|
# with a 422, so the literals must still reach it.
|
||||||
|
body = json.dumps(
|
||||||
|
_payload(
|
||||||
|
provider_id, model_id="aux-rate", pricing=_pricing(**{field: bad})
|
||||||
|
)
|
||||||
|
).encode("utf-8")
|
||||||
|
resp = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers={**_admin_headers(), "content-type": "application/json"},
|
||||||
|
content=body,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 422, field
|
||||||
|
assert (
|
||||||
|
await integration_session.get(ModelRow, ("aux-rate", provider_id)) is None
|
||||||
|
), field
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_admin_single_model_shows_a_non_finite_stored_rate(
|
||||||
|
integration_client: AsyncClient, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""The single-model view answers like the listing it is opened from.
|
||||||
|
|
||||||
|
It is the other view of a row the served-catalog backstop holds back, so it
|
||||||
|
has the same duty to name the rate that needs fixing rather than rendering
|
||||||
|
it as ``null``.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
integration_session.add(
|
||||||
|
ModelRow(
|
||||||
|
id="inf-one",
|
||||||
|
name="inf-one",
|
||||||
|
description="d",
|
||||||
|
created=0,
|
||||||
|
context_length=8192,
|
||||||
|
architecture=json.dumps(
|
||||||
|
{
|
||||||
|
"modality": "text",
|
||||||
|
"input_modalities": ["text"],
|
||||||
|
"output_modalities": ["text"],
|
||||||
|
"tokenizer": "unknown",
|
||||||
|
"instruct_type": None,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
pricing=json.dumps({"prompt": float("inf"), "completion": 2e-06}),
|
||||||
|
upstream_provider_id=provider_id,
|
||||||
|
enabled=True,
|
||||||
|
forwarded_model_id="inf-one",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
resp = await integration_client.get(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models/inf-one",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert resp.status_code == 200
|
||||||
|
assert resp.json()["pricing"]["prompt"] == "inf"
|
||||||
@@ -1,190 +0,0 @@
|
|||||||
import asyncio
|
|
||||||
import secrets
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
from fastapi import HTTPException
|
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
|
||||||
|
|
||||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
|
||||||
from routstr.balance import ChildKeyRequest, create_child_key
|
|
||||||
from routstr.core.db import ApiKey, create_session
|
|
||||||
from routstr.core.settings import settings
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_child_key_flow(integration_session: AsyncSession) -> None:
|
|
||||||
# 1. Create a parent key with balance
|
|
||||||
parent_raw = "parent_test_key_" + secrets.token_hex(4)
|
|
||||||
parent_key = ApiKey(
|
|
||||||
hashed_key=parent_raw,
|
|
||||||
balance=10000, # 10 sats
|
|
||||||
)
|
|
||||||
integration_session.add(parent_key)
|
|
||||||
await integration_session.commit()
|
|
||||||
await integration_session.refresh(parent_key)
|
|
||||||
|
|
||||||
# Mock settings
|
|
||||||
settings.child_key_cost = 1000 # 1 sat
|
|
||||||
|
|
||||||
# 2. Call create_child_key
|
|
||||||
result = await create_child_key(
|
|
||||||
ChildKeyRequest(count=1), parent_key, integration_session
|
|
||||||
)
|
|
||||||
|
|
||||||
assert "api_keys" in result
|
|
||||||
assert result["cost_msats"] == 1000
|
|
||||||
assert result["parent_balance"] == 9000
|
|
||||||
|
|
||||||
child_key_raw = result["api_keys"][0][3:] # remove sk-
|
|
||||||
|
|
||||||
# 3. Verify child key exists in DB
|
|
||||||
child_key_db = await integration_session.get(ApiKey, child_key_raw)
|
|
||||||
assert child_key_db is not None
|
|
||||||
assert child_key_db.parent_key_hash == parent_key.hashed_key
|
|
||||||
assert child_key_db.balance == 0
|
|
||||||
|
|
||||||
# 4. Test payment with child key
|
|
||||||
cost = 500
|
|
||||||
await pay_for_request(child_key_db, cost, integration_session)
|
|
||||||
|
|
||||||
# Refresh keys
|
|
||||||
await integration_session.refresh(parent_key)
|
|
||||||
await integration_session.refresh(child_key_db)
|
|
||||||
|
|
||||||
# Parent should be charged
|
|
||||||
assert parent_key.reserved_balance == 500
|
|
||||||
assert parent_key.total_requests == 1
|
|
||||||
|
|
||||||
# Child should have total_requests incremented
|
|
||||||
assert child_key_db.total_requests == 1
|
|
||||||
|
|
||||||
# 5. Test adjustment
|
|
||||||
response_data = {"model": "test-model", "usage": {"total_tokens": 10}}
|
|
||||||
|
|
||||||
# Mock calculate_cost
|
|
||||||
import routstr.auth
|
|
||||||
from routstr.payment.cost_calculation import CostData
|
|
||||||
|
|
||||||
async def mock_calculate_cost(*args: Any, **kwargs: Any) -> CostData:
|
|
||||||
return CostData(
|
|
||||||
base_msats=0, input_msats=200, output_msats=200, total_msats=400
|
|
||||||
)
|
|
||||||
|
|
||||||
# Patch calculate_cost
|
|
||||||
original_calculate_cost = routstr.auth.calculate_cost
|
|
||||||
routstr.auth.calculate_cost = mock_calculate_cost
|
|
||||||
|
|
||||||
try:
|
|
||||||
adjustment = await adjust_payment_for_tokens(
|
|
||||||
child_key_db, response_data, integration_session, 500, None, None
|
|
||||||
)
|
|
||||||
assert adjustment["total_msats"] == 400
|
|
||||||
|
|
||||||
# Refresh keys
|
|
||||||
await integration_session.refresh(parent_key)
|
|
||||||
await integration_session.refresh(child_key_db)
|
|
||||||
|
|
||||||
# Parent should have updated balance and total_spent
|
|
||||||
assert parent_key.reserved_balance == 0
|
|
||||||
assert parent_key.balance == 9000 - 400
|
|
||||||
assert (
|
|
||||||
parent_key.total_spent == 1400
|
|
||||||
) # 1000 for child key creation + 400 for request
|
|
||||||
|
|
||||||
# Child should also have total_spent updated
|
|
||||||
assert child_key_db.total_spent == 400
|
|
||||||
|
|
||||||
finally:
|
|
||||||
routstr.auth.calculate_cost = original_calculate_cost
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_child_key_insufficient_balance(
|
|
||||||
integration_session: AsyncSession,
|
|
||||||
) -> None:
|
|
||||||
parent_key = ApiKey(
|
|
||||||
hashed_key="poor_parent_" + secrets.token_hex(4),
|
|
||||||
balance=500,
|
|
||||||
)
|
|
||||||
integration_session.add(parent_key)
|
|
||||||
await integration_session.commit()
|
|
||||||
await integration_session.refresh(parent_key)
|
|
||||||
|
|
||||||
settings.child_key_cost = 1000
|
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc:
|
|
||||||
await create_child_key(
|
|
||||||
ChildKeyRequest(count=1), parent_key, integration_session
|
|
||||||
)
|
|
||||||
assert exc.value.status_code == 402
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_concurrent_child_key_creation_is_atomic(
|
|
||||||
patched_db_engine: None,
|
|
||||||
) -> None:
|
|
||||||
"""Two concurrent create_child_key() calls with balance for exactly one must
|
|
||||||
result in exactly one success and one 402, with the parent balance deducted
|
|
||||||
only once."""
|
|
||||||
child_key_cost = 1000
|
|
||||||
settings.child_key_cost = child_key_cost
|
|
||||||
|
|
||||||
parent_hash = f"parent_concurrent_{secrets.token_hex(8)}"
|
|
||||||
async with create_session() as session:
|
|
||||||
parent = ApiKey(hashed_key=parent_hash, balance=child_key_cost)
|
|
||||||
session.add(parent)
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
results: list[str] = []
|
|
||||||
|
|
||||||
async def attempt() -> None:
|
|
||||||
async with create_session() as session:
|
|
||||||
fresh_parent = await session.get(ApiKey, parent_hash)
|
|
||||||
assert fresh_parent is not None
|
|
||||||
try:
|
|
||||||
await create_child_key(ChildKeyRequest(count=1), fresh_parent, session)
|
|
||||||
results.append("success")
|
|
||||||
except HTTPException as exc:
|
|
||||||
assert exc.status_code == 402
|
|
||||||
results.append("blocked")
|
|
||||||
|
|
||||||
await asyncio.gather(attempt(), attempt())
|
|
||||||
|
|
||||||
assert sorted(results) == ["blocked", "success"], (
|
|
||||||
f"Expected exactly one success and one 402, got: {results}"
|
|
||||||
)
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
final = await session.get(ApiKey, parent_hash)
|
|
||||||
assert final is not None
|
|
||||||
|
|
||||||
assert final.balance == 0, (
|
|
||||||
f"Balance should be fully deducted once: expected 0, got {final.balance}"
|
|
||||||
)
|
|
||||||
assert final.total_spent == child_key_cost, (
|
|
||||||
f"total_spent should equal one deduction: expected {child_key_cost}, "
|
|
||||||
f"got {final.total_spent}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_child_key_cannot_create_child(integration_session: AsyncSession) -> None:
|
|
||||||
parent_key = ApiKey(
|
|
||||||
hashed_key="parent_" + secrets.token_hex(4),
|
|
||||||
balance=10000,
|
|
||||||
)
|
|
||||||
child_key = ApiKey(
|
|
||||||
hashed_key="child_" + secrets.token_hex(4),
|
|
||||||
balance=0,
|
|
||||||
parent_key_hash=parent_key.hashed_key,
|
|
||||||
)
|
|
||||||
integration_session.add(parent_key)
|
|
||||||
integration_session.add(child_key)
|
|
||||||
await integration_session.commit()
|
|
||||||
await integration_session.refresh(child_key)
|
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc:
|
|
||||||
await create_child_key(ChildKeyRequest(count=1), child_key, integration_session)
|
|
||||||
assert exc.value.status_code == 400
|
|
||||||
assert "Cannot create a child key for another child key" in str(exc.value.detail)
|
|
||||||
@@ -1,102 +0,0 @@
|
|||||||
from typing import Any
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
from httpx import AsyncClient
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_wallet_info_returns_child_keys(
|
|
||||||
integration_client: AsyncClient,
|
|
||||||
authenticated_client: AsyncClient,
|
|
||||||
integration_session: Any,
|
|
||||||
) -> None:
|
|
||||||
"""Test that GET /v1/wallet/info returns child keys for a parent key"""
|
|
||||||
|
|
||||||
# 1. Get parent info to find its hashed_key
|
|
||||||
response = await authenticated_client.get("/v1/wallet/info")
|
|
||||||
assert response.status_code == 200
|
|
||||||
parent_data = response.json()
|
|
||||||
parent_data["api_key"]
|
|
||||||
|
|
||||||
# 2. Create child keys for this parent
|
|
||||||
# We need to use the parent's authentication for this
|
|
||||||
child_payload = {"count": 2, "balance_limit": 1000, "balance_limit_reset": "daily"}
|
|
||||||
create_response = await authenticated_client.post(
|
|
||||||
"/v1/wallet/child-key", json=child_payload
|
|
||||||
)
|
|
||||||
assert create_response.status_code == 200
|
|
||||||
create_data = create_response.json()
|
|
||||||
child_keys = create_data["api_keys"]
|
|
||||||
assert len(child_keys) == 2
|
|
||||||
|
|
||||||
# 3. Call /info again and check for child_keys
|
|
||||||
info_response = await authenticated_client.get("/v1/wallet/info")
|
|
||||||
assert info_response.status_code == 200
|
|
||||||
info_data = info_response.json()
|
|
||||||
|
|
||||||
assert "child_keys" in info_data
|
|
||||||
assert len(info_data["child_keys"]) == 2
|
|
||||||
|
|
||||||
# Verify child key details
|
|
||||||
for ck in info_data["child_keys"]:
|
|
||||||
assert ck["api_key"] in child_keys
|
|
||||||
assert ck["balance_limit"] == 1000
|
|
||||||
assert ck["balance_limit_reset"] == "daily"
|
|
||||||
assert "total_spent" in ck
|
|
||||||
assert "total_requests" in ck
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_wallet_info_child_key_no_child_keys(
|
|
||||||
integration_client: AsyncClient,
|
|
||||||
authenticated_client: AsyncClient,
|
|
||||||
integration_session: Any,
|
|
||||||
) -> None:
|
|
||||||
"""Test that GET /v1/wallet/info for a child key does NOT return child_keys"""
|
|
||||||
|
|
||||||
# 1. Create a child key
|
|
||||||
child_payload = {"count": 1}
|
|
||||||
create_response = await authenticated_client.post(
|
|
||||||
"/v1/wallet/child-key", json=child_payload
|
|
||||||
)
|
|
||||||
assert create_response.status_code == 200
|
|
||||||
child_key = create_response.json()["api_keys"][0]
|
|
||||||
|
|
||||||
# 2. Use the child key to get its info
|
|
||||||
integration_client.headers["Authorization"] = f"Bearer {child_key}"
|
|
||||||
info_response = await integration_client.get("/v1/wallet/info")
|
|
||||||
assert info_response.status_code == 200
|
|
||||||
info_data = info_response.json()
|
|
||||||
parent_key = authenticated_client._test_api_key # type: ignore[attr-defined]
|
|
||||||
parent_key_hash = parent_key.removeprefix("sk-")
|
|
||||||
|
|
||||||
assert info_data["is_child"] is True
|
|
||||||
assert "child_keys" not in info_data
|
|
||||||
assert "parent_key" not in info_data
|
|
||||||
assert info_data["parent_key_preview"] == parent_key_hash[:8] + "..."
|
|
||||||
assert info_data["parent_key_preview"] not in {parent_key, parent_key_hash}
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_account_info_root_returns_child_keys(
|
|
||||||
authenticated_client: AsyncClient,
|
|
||||||
) -> None:
|
|
||||||
"""Test that GET / returns child keys for a parent key (root endpoint)"""
|
|
||||||
|
|
||||||
# 1. Create a child key
|
|
||||||
child_payload = {"count": 1}
|
|
||||||
await authenticated_client.post("/v1/wallet/child-key", json=child_payload)
|
|
||||||
|
|
||||||
# 2. Call root endpoint /v1/balance/
|
|
||||||
# Note: routstr/balance.py defines router = APIRouter()
|
|
||||||
# and it is included in balance_router with prefix /v1/balance
|
|
||||||
# The endpoint is @router.get("/")
|
|
||||||
response = await authenticated_client.get("/v1/balance/")
|
|
||||||
assert response.status_code == 200
|
|
||||||
data = response.json()
|
|
||||||
|
|
||||||
assert "child_keys" in data
|
|
||||||
assert len(data["child_keys"]) >= 1
|
|
||||||
@@ -31,7 +31,7 @@ class TestNetworkFailureScenarios:
|
|||||||
AsyncMock(side_effect=ConnectError("Mint service unavailable")),
|
AsyncMock(side_effect=ConnectError("Mint service unavailable")),
|
||||||
),
|
),
|
||||||
patch(
|
patch(
|
||||||
"routstr.balance.send_token",
|
"routstr.refund.send_token",
|
||||||
AsyncMock(side_effect=ConnectError("Mint service unavailable")),
|
AsyncMock(side_effect=ConnectError("Mint service unavailable")),
|
||||||
),
|
),
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from httpx import AsyncClient
|
|||||||
from sqlmodel import select
|
from sqlmodel import select
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from routstr.core.db import ApiKey, ReservationRelease
|
from routstr.core.db import ReservationRelease
|
||||||
from routstr.payment.models import Architecture, Model, Pricing
|
from routstr.payment.models import Architecture, Model, Pricing
|
||||||
from routstr.proxy import refresh_model_maps
|
from routstr.proxy import refresh_model_maps
|
||||||
from routstr.upstream.base import BaseUpstreamProvider
|
from routstr.upstream.base import BaseUpstreamProvider
|
||||||
@@ -531,103 +531,6 @@ async def test_failover_beyond_balance_envelope_is_rejected(
|
|||||||
# fallback must be rejected before its upstream is ever contacted.
|
# fallback must be rejected before its upstream is ever contacted.
|
||||||
assert response.status_code == 402
|
assert response.status_code == 402
|
||||||
assert [r.url.host for r in sent_requests] == ["cheap.example.com"]
|
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
|
@pytest.fixture
|
||||||
async def raised_envelope_provider_maps(
|
async def raised_envelope_provider_maps(
|
||||||
patched_db_engine: None,
|
patched_db_engine: None,
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from unittest.mock import patch
|
|||||||
import pytest
|
import pytest
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from routstr.core.db import ApiKey
|
from routstr.core.db import ApiKey, ReservationRelease
|
||||||
from routstr.payment.cost_calculation import CostData
|
from routstr.payment.cost_calculation import CostData
|
||||||
|
|
||||||
|
|
||||||
@@ -34,20 +34,30 @@ def _cost_data(total_msats: int) -> CostData:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_overrun_charges_after_reservation_swept(
|
async def test_overrun_with_corrupted_aggregate_releases_without_charging(
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Overrun finalize must charge even when the reservation was already released."""
|
"""An overrun whose aggregate reservation was externally zeroed must not
|
||||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
charge: subtracting the reservation would have to clamp, which can erase
|
||||||
|
sibling reservations. The reservation is released and the charge dropped.
|
||||||
|
"""
|
||||||
|
from routstr import auth
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
deducted_max_cost = 990 # discounted reservation
|
deducted_max_cost = 990 # discounted reservation
|
||||||
actual_token_cost = 1000 # actual cost overruns the reservation
|
actual_token_cost = 1000 # actual cost overruns the reservation
|
||||||
|
|
||||||
# Sweeper has zeroed reserved_balance but left balance untouched.
|
# Something zeroed reserved_balance under an active durable reservation.
|
||||||
key = _make_key(balance=1000, reserved=0)
|
key = _make_key(balance=1000, reserved=0)
|
||||||
|
key_hash = key.hashed_key
|
||||||
integration_session.add(key)
|
integration_session.add(key)
|
||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
await pay_for_request(key, deducted_max_cost, integration_session)
|
await pay_for_request(key, deducted_max_cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
key.reserved_balance = 0
|
key.reserved_balance = 0
|
||||||
integration_session.add(key)
|
integration_session.add(key)
|
||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
@@ -61,20 +71,67 @@ async def test_overrun_charges_after_reservation_swept(
|
|||||||
"routstr.auth.calculate_cost",
|
"routstr.auth.calculate_cost",
|
||||||
return_value=_cost_data(actual_token_cost),
|
return_value=_cost_data(actual_token_cost),
|
||||||
):
|
):
|
||||||
await adjust_payment_for_tokens(
|
result = await adjust_payment_for_tokens(
|
||||||
key, response_data, integration_session, deducted_max_cost, None, None
|
key, response_data, integration_session, deducted_max_cost, None, None
|
||||||
)
|
)
|
||||||
|
|
||||||
await integration_session.refresh(key)
|
assert result["charged_msats"] == 0
|
||||||
|
integration_session.expunge_all()
|
||||||
|
key_row = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key_row is not None
|
||||||
|
|
||||||
assert key.total_spent == actual_token_cost, (
|
assert key_row.total_spent == 0, "corrupted aggregate must not be charged into"
|
||||||
f"Request was not billed (total_spent={key.total_spent}) — free response bug"
|
assert key_row.balance == 1000
|
||||||
|
assert key_row.reserved_balance == 0
|
||||||
|
|
||||||
|
# The corrupt reservation must reach a terminal state — an active leftover
|
||||||
|
# would be renewed by its heartbeat forever and poison stale cleanup.
|
||||||
|
record = await integration_session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
assert reservation.release_id not in auth._reservation_heartbeats
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_missing_usage_never_turns_reservation_into_charge(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A reservation is an authorization ceiling, not evidence of usage.
|
||||||
|
|
||||||
|
Upstream handlers should provide locally estimated usage when possible. If
|
||||||
|
no measurement or estimate reaches settlement, release the reservation
|
||||||
|
rather than charging its full value.
|
||||||
|
"""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
)
|
)
|
||||||
assert key.balance == 1000 - actual_token_cost, (
|
|
||||||
f"Balance not charged: {key.balance}"
|
reserved = 4_000
|
||||||
|
key = _make_key(balance=10_000, reserved=0)
|
||||||
|
key_hash = key.hashed_key
|
||||||
|
integration_session.add(key)
|
||||||
|
await integration_session.commit()
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
# No `usage` key at all — the upstream stream dropped its final usage chunk.
|
||||||
|
response_data = {"model": "test-model"}
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
response_data,
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
)
|
)
|
||||||
assert key.balance >= 0
|
|
||||||
assert key.reserved_balance == 0
|
assert result["charged_msats"] == 0
|
||||||
|
integration_session.expunge_all()
|
||||||
|
key_row = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key_row is not None
|
||||||
|
assert key_row.total_spent == 0
|
||||||
|
assert key_row.balance == 10_000
|
||||||
|
assert key_row.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@@ -1,14 +1,10 @@
|
|||||||
import asyncio
|
|
||||||
import time
|
import time
|
||||||
from datetime import datetime, timedelta
|
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from fastapi import HTTPException
|
|
||||||
from sqlmodel import select
|
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
from routstr.auth import pay_for_request
|
from routstr.auth import pay_for_request
|
||||||
from routstr.core.db import ApiKey, create_session
|
from routstr.core.db import ApiKey
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -25,367 +21,6 @@ async def test_key_validity_date(integration_session: AsyncSession) -> None:
|
|||||||
assert "expired" in str(excinfo.value).lower()
|
assert "expired" in str(excinfo.value).lower()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_key_balance_limit(integration_session: AsyncSession) -> None:
|
|
||||||
# 1. Create a key with a balance limit
|
|
||||||
key = ApiKey(
|
|
||||||
hashed_key="limited_key", balance=10000, balance_limit=500, total_spent=450
|
|
||||||
)
|
|
||||||
integration_session.add(key)
|
|
||||||
await integration_session.commit()
|
|
||||||
|
|
||||||
# 2. Try to pay for a request that exceeds the limit
|
|
||||||
with pytest.raises(Exception) as excinfo:
|
|
||||||
await pay_for_request(key, 100, integration_session)
|
|
||||||
assert "limit exceeded" in str(excinfo.value).lower()
|
|
||||||
|
|
||||||
# 3. Try to pay for a request that fits
|
|
||||||
await pay_for_request(key, 50, integration_session)
|
|
||||||
await integration_session.refresh(key)
|
|
||||||
# Note: total_spent is updated in adjust_payment_for_tokens,
|
|
||||||
# but pay_for_request checks it.
|
|
||||||
# In our current logic, pay_for_request checks (total_spent + cost) > balance_limit.
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_key_daily_reset_policy(integration_session: AsyncSession) -> None:
|
|
||||||
# 1. Create a key with a daily reset policy and old reset date
|
|
||||||
yesterday = int((datetime.now() - timedelta(days=1)).timestamp())
|
|
||||||
key = ApiKey(
|
|
||||||
hashed_key="daily_reset_key",
|
|
||||||
balance=10000,
|
|
||||||
balance_limit=1000,
|
|
||||||
balance_limit_reset="daily",
|
|
||||||
balance_limit_reset_date=yesterday,
|
|
||||||
total_spent=900,
|
|
||||||
)
|
|
||||||
integration_session.add(key)
|
|
||||||
await integration_session.commit()
|
|
||||||
|
|
||||||
# 2. Pay for a request - should trigger reset first because it's a new day
|
|
||||||
# Request is 200, total_spent is 900. 900+200 > 1000,
|
|
||||||
# but reset should happen making total_spent 0, then 0+200 < 1000.
|
|
||||||
await pay_for_request(key, 200, integration_session)
|
|
||||||
|
|
||||||
await integration_session.refresh(key)
|
|
||||||
assert key.total_spent == 0 # Reset in pay_for_request happens before charging
|
|
||||||
# Wait, the charging logic in pay_for_request increments parent/billing_key's total_requests,
|
|
||||||
# but total_spent is updated in adjust_payment_for_tokens.
|
|
||||||
# However, the reset logic sets total_spent to 0.
|
|
||||||
assert key.balance_limit_reset_date is not None
|
|
||||||
assert key.balance_limit_reset_date > yesterday
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_periodic_key_reset_job(integration_session: AsyncSession) -> None:
|
|
||||||
# 1. Create multiple keys needing reset
|
|
||||||
yesterday = int((datetime.now() - timedelta(days=1)).timestamp())
|
|
||||||
key1 = ApiKey(
|
|
||||||
hashed_key="job_reset_key_1",
|
|
||||||
balance=1000,
|
|
||||||
balance_limit=1000,
|
|
||||||
balance_limit_reset="daily",
|
|
||||||
balance_limit_reset_date=yesterday,
|
|
||||||
total_spent=500,
|
|
||||||
)
|
|
||||||
key2 = ApiKey(
|
|
||||||
hashed_key="job_reset_key_2",
|
|
||||||
balance=1000,
|
|
||||||
balance_limit=1000,
|
|
||||||
balance_limit_reset="daily",
|
|
||||||
balance_limit_reset_date=yesterday,
|
|
||||||
total_spent=800,
|
|
||||||
)
|
|
||||||
integration_session.add(key1)
|
|
||||||
integration_session.add(key2)
|
|
||||||
await integration_session.commit()
|
|
||||||
|
|
||||||
# 2. Run the periodic reset logic manually (mocking the background task loop)
|
|
||||||
# We can't easily run the actual loop because it has a sleep,
|
|
||||||
# but we can test the logic inside.
|
|
||||||
|
|
||||||
# Implementation of periodic_key_reset logic for testing:
|
|
||||||
stmt = select(ApiKey).where(ApiKey.balance_limit_reset != None) # noqa: E711
|
|
||||||
keys = (await integration_session.exec(stmt)).all()
|
|
||||||
now = int(time.time())
|
|
||||||
for k in keys:
|
|
||||||
if k.hashed_key in ["job_reset_key_1", "job_reset_key_2"]:
|
|
||||||
k.total_spent = 0
|
|
||||||
k.balance_limit_reset_date = now
|
|
||||||
integration_session.add(k)
|
|
||||||
await integration_session.commit()
|
|
||||||
|
|
||||||
# 3. Verify resets
|
|
||||||
await integration_session.refresh(key1)
|
|
||||||
await integration_session.refresh(key2)
|
|
||||||
assert key1.total_spent == 0
|
|
||||||
assert key2.total_spent == 0
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_balance_limit_enforced_atomically_under_concurrency(
|
|
||||||
patched_db_engine: None,
|
|
||||||
) -> None:
|
|
||||||
parent_hash = "parent_limit_atomic"
|
|
||||||
child_hash = "child_limit_atomic"
|
|
||||||
cost = 300
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
parent = ApiKey(hashed_key=parent_hash, balance=10000)
|
|
||||||
child = ApiKey(
|
|
||||||
hashed_key=child_hash,
|
|
||||||
balance=0,
|
|
||||||
parent_key_hash=parent_hash,
|
|
||||||
balance_limit=cost,
|
|
||||||
total_spent=0,
|
|
||||||
)
|
|
||||||
session.add(parent)
|
|
||||||
session.add(child)
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
results: list[str] = []
|
|
||||||
|
|
||||||
async def attempt() -> None:
|
|
||||||
async with create_session() as session:
|
|
||||||
fresh_child = await session.get(ApiKey, child_hash)
|
|
||||||
assert fresh_child is not None
|
|
||||||
try:
|
|
||||||
await pay_for_request(fresh_child, cost, session)
|
|
||||||
results.append("success")
|
|
||||||
except HTTPException as exc:
|
|
||||||
assert exc.status_code == 402
|
|
||||||
results.append("blocked")
|
|
||||||
|
|
||||||
await asyncio.gather(attempt(), attempt())
|
|
||||||
|
|
||||||
assert sorted(results) == ["blocked", "success"], (
|
|
||||||
f"Expected exactly one success and one 402, got: {results}"
|
|
||||||
)
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
final_child = await session.get(ApiKey, child_hash)
|
|
||||||
assert final_child is not None
|
|
||||||
|
|
||||||
assert final_child.reserved_balance == cost, (
|
|
||||||
f"Child reserved_balance should equal one reservation, "
|
|
||||||
f"got {final_child.reserved_balance}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_parallel_payments_with_parent_and_child_key(
|
|
||||||
patched_db_engine: None,
|
|
||||||
) -> None:
|
|
||||||
parent_hash = "parent_parallel_mixed"
|
|
||||||
child_hash = "child_parallel_mixed"
|
|
||||||
cost = 300
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
parent = ApiKey(hashed_key=parent_hash, balance=10000)
|
|
||||||
child = ApiKey(
|
|
||||||
hashed_key=child_hash,
|
|
||||||
balance=0,
|
|
||||||
parent_key_hash=parent_hash,
|
|
||||||
balance_limit=2 * cost,
|
|
||||||
)
|
|
||||||
session.add(parent)
|
|
||||||
session.add(child)
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
async def attempt(key_hash: str) -> str:
|
|
||||||
async with create_session() as session:
|
|
||||||
fresh_key = await session.get(ApiKey, key_hash)
|
|
||||||
assert fresh_key is not None
|
|
||||||
try:
|
|
||||||
await pay_for_request(fresh_key, cost, session)
|
|
||||||
return "success"
|
|
||||||
except HTTPException as exc:
|
|
||||||
assert exc.status_code == 402
|
|
||||||
return "blocked"
|
|
||||||
|
|
||||||
results = await asyncio.gather(attempt(parent_hash), attempt(child_hash))
|
|
||||||
|
|
||||||
assert results == ["success", "success"], (
|
|
||||||
f"Both parent and child payments should succeed, got: {results}"
|
|
||||||
)
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
final_parent = await session.get(ApiKey, parent_hash)
|
|
||||||
final_child = await session.get(ApiKey, child_hash)
|
|
||||||
assert final_parent is not None
|
|
||||||
assert final_child is not None
|
|
||||||
|
|
||||||
# Both requests bill the parent; only the child request reserves on the child.
|
|
||||||
assert final_parent.reserved_balance == 2 * cost
|
|
||||||
assert final_parent.total_requests == 2
|
|
||||||
assert final_child.reserved_balance == cost
|
|
||||||
assert final_child.total_requests == 1
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_balance_limit_with_existing_total_spent_under_concurrency(
|
|
||||||
patched_db_engine: None,
|
|
||||||
) -> None:
|
|
||||||
parent_hash = "parent_total_spent"
|
|
||||||
child_hash = "child_total_spent"
|
|
||||||
cost = 300
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
parent = ApiKey(hashed_key=parent_hash, balance=10000)
|
|
||||||
# 700 already spent against a 1000 limit: only one more 300 request fits.
|
|
||||||
child = ApiKey(
|
|
||||||
hashed_key=child_hash,
|
|
||||||
balance=0,
|
|
||||||
parent_key_hash=parent_hash,
|
|
||||||
balance_limit=1000,
|
|
||||||
total_spent=700,
|
|
||||||
)
|
|
||||||
session.add(parent)
|
|
||||||
session.add(child)
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
results: list[str] = []
|
|
||||||
|
|
||||||
async def attempt() -> None:
|
|
||||||
async with create_session() as session:
|
|
||||||
fresh_child = await session.get(ApiKey, child_hash)
|
|
||||||
assert fresh_child is not None
|
|
||||||
try:
|
|
||||||
await pay_for_request(fresh_child, cost, session)
|
|
||||||
results.append("success")
|
|
||||||
except HTTPException as exc:
|
|
||||||
assert exc.status_code == 402
|
|
||||||
results.append("blocked")
|
|
||||||
|
|
||||||
await asyncio.gather(attempt(), attempt())
|
|
||||||
|
|
||||||
assert sorted(results) == ["blocked", "success"], (
|
|
||||||
f"Expected exactly one success and one 402, got: {results}"
|
|
||||||
)
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
final_child = await session.get(ApiKey, child_hash)
|
|
||||||
assert final_child is not None
|
|
||||||
|
|
||||||
assert final_child.reserved_balance == cost
|
|
||||||
assert final_child.total_spent == 700
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_balance_limit_with_existing_reserved_balance(
|
|
||||||
patched_db_engine: None,
|
|
||||||
) -> None:
|
|
||||||
parent_hash = "parent_reserved_set"
|
|
||||||
blocked_hash = "child_reserved_blocked"
|
|
||||||
allowed_hash = "child_reserved_allowed"
|
|
||||||
cost = 300
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
parent = ApiKey(hashed_key=parent_hash, balance=10000)
|
|
||||||
# 800 already reserved against a 1000 limit: another 300 must be rejected.
|
|
||||||
blocked_child = ApiKey(
|
|
||||||
hashed_key=blocked_hash,
|
|
||||||
balance=0,
|
|
||||||
parent_key_hash=parent_hash,
|
|
||||||
balance_limit=1000,
|
|
||||||
reserved_balance=800,
|
|
||||||
)
|
|
||||||
# 500 reserved against a 1000 limit: another 300 still fits.
|
|
||||||
allowed_child = ApiKey(
|
|
||||||
hashed_key=allowed_hash,
|
|
||||||
balance=0,
|
|
||||||
parent_key_hash=parent_hash,
|
|
||||||
balance_limit=1000,
|
|
||||||
reserved_balance=500,
|
|
||||||
)
|
|
||||||
session.add(parent)
|
|
||||||
session.add(blocked_child)
|
|
||||||
session.add(allowed_child)
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
fresh_blocked = await session.get(ApiKey, blocked_hash)
|
|
||||||
assert fresh_blocked is not None
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
|
||||||
await pay_for_request(fresh_blocked, cost, session)
|
|
||||||
assert exc_info.value.status_code == 402
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
fresh_allowed = await session.get(ApiKey, allowed_hash)
|
|
||||||
assert fresh_allowed is not None
|
|
||||||
await pay_for_request(fresh_allowed, cost, session)
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
final_blocked = await session.get(ApiKey, blocked_hash)
|
|
||||||
final_allowed = await session.get(ApiKey, allowed_hash)
|
|
||||||
final_parent = await session.get(ApiKey, parent_hash)
|
|
||||||
assert final_blocked is not None
|
|
||||||
assert final_allowed is not None
|
|
||||||
assert final_parent is not None
|
|
||||||
|
|
||||||
assert final_blocked.reserved_balance == 800, "Rejected request must not reserve"
|
|
||||||
assert final_blocked.total_requests == 0
|
|
||||||
assert final_allowed.reserved_balance == 500 + cost
|
|
||||||
assert final_allowed.total_requests == 1
|
|
||||||
# Only the allowed request should have billed the parent.
|
|
||||||
assert final_parent.reserved_balance == cost
|
|
||||||
assert final_parent.total_requests == 1
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_child_reservation_discarded_when_parent_balance_depleted(
|
|
||||||
patched_db_engine: None,
|
|
||||||
) -> None:
|
|
||||||
parent_hash = "parent_depleted"
|
|
||||||
child_hash = "child_depleted"
|
|
||||||
cost = 300
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
# Parent can only afford one request; child has no balance_limit.
|
|
||||||
parent = ApiKey(hashed_key=parent_hash, balance=cost)
|
|
||||||
child = ApiKey(
|
|
||||||
hashed_key=child_hash,
|
|
||||||
balance=0,
|
|
||||||
parent_key_hash=parent_hash,
|
|
||||||
)
|
|
||||||
session.add(parent)
|
|
||||||
session.add(child)
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
results: list[str] = []
|
|
||||||
|
|
||||||
async def attempt() -> None:
|
|
||||||
async with create_session() as session:
|
|
||||||
fresh_child = await session.get(ApiKey, child_hash)
|
|
||||||
assert fresh_child is not None
|
|
||||||
try:
|
|
||||||
await pay_for_request(fresh_child, cost, session)
|
|
||||||
results.append("success")
|
|
||||||
except HTTPException as exc:
|
|
||||||
assert exc.status_code == 402
|
|
||||||
results.append("blocked")
|
|
||||||
|
|
||||||
await asyncio.gather(attempt(), attempt())
|
|
||||||
|
|
||||||
assert sorted(results) == ["blocked", "success"], (
|
|
||||||
f"Expected exactly one success and one 402, got: {results}"
|
|
||||||
)
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
final_parent = await session.get(ApiKey, parent_hash)
|
|
||||||
final_child = await session.get(ApiKey, child_hash)
|
|
||||||
assert final_parent is not None
|
|
||||||
assert final_child is not None
|
|
||||||
|
|
||||||
assert final_parent.reserved_balance == cost
|
|
||||||
# The failed request must not leave a committed reservation on the child.
|
|
||||||
assert final_child.reserved_balance == cost, (
|
|
||||||
f"Child reserved_balance should reflect only the successful request, "
|
|
||||||
f"got {final_child.reserved_balance}"
|
|
||||||
)
|
|
||||||
assert final_child.total_requests == 1
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_refund_does_not_delete_key(integration_session: AsyncSession) -> None:
|
async def test_refund_does_not_delete_key(integration_session: AsyncSession) -> None:
|
||||||
# This requires mocking the router call or testing the logic in balance.py
|
# This requires mocking the router call or testing the logic in balance.py
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
"""Integration tests for Lightning invoice key constraint fields.
|
"""Integration tests for Lightning invoice key constraint fields.
|
||||||
|
|
||||||
Covers two things:
|
Covers two things:
|
||||||
- The three constraint fields (balance_limit, balance_limit_reset, validity_date)
|
- The validity_date constraint field is persisted on LightningInvoice and
|
||||||
are persisted on LightningInvoice and survive a DB round-trip.
|
survives a DB round-trip.
|
||||||
- The production-path API-key record helper propagates those fields to the
|
- The production-path API-key record helper propagates it to the created
|
||||||
created ApiKey, so the constraints are actually enforced when the key is used.
|
ApiKey, so the constraint is actually enforced when the key is used.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -67,32 +67,6 @@ def mock_wallet_mint() -> object:
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_invoice_persists_balance_limit(
|
|
||||||
integration_session: AsyncSession,
|
|
||||||
) -> None:
|
|
||||||
invoice = _make_invoice(balance_limit=5000)
|
|
||||||
integration_session.add(invoice)
|
|
||||||
await integration_session.commit()
|
|
||||||
|
|
||||||
stored = await integration_session.get(LightningInvoice, invoice.id)
|
|
||||||
assert stored is not None
|
|
||||||
assert stored.balance_limit == 5000
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_invoice_persists_balance_limit_reset(
|
|
||||||
integration_session: AsyncSession,
|
|
||||||
) -> None:
|
|
||||||
invoice = _make_invoice(balance_limit=5000, balance_limit_reset="daily")
|
|
||||||
integration_session.add(invoice)
|
|
||||||
await integration_session.commit()
|
|
||||||
|
|
||||||
stored = await integration_session.get(LightningInvoice, invoice.id)
|
|
||||||
assert stored is not None
|
|
||||||
assert stored.balance_limit_reset == "daily"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_invoice_persists_validity_date(
|
async def test_invoice_persists_validity_date(
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
@@ -112,38 +86,6 @@ async def test_invoice_persists_validity_date(
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_created_key_receives_balance_limit(
|
|
||||||
integration_session: AsyncSession,
|
|
||||||
) -> None:
|
|
||||||
invoice = _make_invoice(balance_limit=8000)
|
|
||||||
integration_session.add(invoice)
|
|
||||||
await integration_session.flush()
|
|
||||||
|
|
||||||
api_key = await _create_api_key_record(invoice, integration_session)
|
|
||||||
await integration_session.commit()
|
|
||||||
|
|
||||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
|
||||||
assert stored_key is not None
|
|
||||||
assert stored_key.balance_limit == 8000
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_created_key_receives_balance_limit_reset(
|
|
||||||
integration_session: AsyncSession,
|
|
||||||
) -> None:
|
|
||||||
invoice = _make_invoice(balance_limit=8000, balance_limit_reset="monthly")
|
|
||||||
integration_session.add(invoice)
|
|
||||||
await integration_session.flush()
|
|
||||||
|
|
||||||
api_key = await _create_api_key_record(invoice, integration_session)
|
|
||||||
await integration_session.commit()
|
|
||||||
|
|
||||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
|
||||||
assert stored_key is not None
|
|
||||||
assert stored_key.balance_limit_reset == "monthly"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_created_key_receives_validity_date(
|
async def test_created_key_receives_validity_date(
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
@@ -422,8 +364,6 @@ async def test_created_key_without_constraints_has_none_fields(
|
|||||||
|
|
||||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||||
assert stored_key is not None
|
assert stored_key is not None
|
||||||
assert stored_key.balance_limit is None
|
|
||||||
assert stored_key.balance_limit_reset is None
|
|
||||||
assert stored_key.validity_date is None
|
assert stored_key.validity_date is None
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,221 @@
|
|||||||
|
"""Cover that an admin price write reaches the served catalogue (GET /v1/models).
|
||||||
|
|
||||||
|
A price edit changes the served price; the served price is fee-adjusted while the
|
||||||
|
admin read-back is raw; a disabled model leaves the catalogue but keeps its row.
|
||||||
|
|
||||||
|
Each test reaches the served model by id rather than iterating ``data["data"]``,
|
||||||
|
so an empty catalogue fails these tests instead of skipping them.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Iterator
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from httpx import AsyncClient
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from routstr.core.admin import admin_sessions
|
||||||
|
from routstr.core.db import ModelRow, UpstreamProviderRow
|
||||||
|
from routstr.proxy import reinitialize_upstreams
|
||||||
|
|
||||||
|
|
||||||
|
# The conftest patches ``routstr.payment.price.sats_usd_price``, but
|
||||||
|
# ``models.py`` imports it as ``from .price import sats_usd_price`` — a
|
||||||
|
# local binding the conftest-level patch cannot reach. Pin it here so
|
||||||
|
# every test that goes through ``_row_to_model`` gets a real sats price.
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _pin_sats_usd() -> Iterator[None]:
|
||||||
|
with patch("routstr.payment.models.sats_usd_price", return_value=0.0005):
|
||||||
|
yield
|
||||||
|
|
||||||
|
|
||||||
|
def _admin_headers() -> dict[str, str]:
|
||||||
|
token = "test-propagation-token"
|
||||||
|
admin_sessions[token] = int(
|
||||||
|
(datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp()
|
||||||
|
)
|
||||||
|
return {"Authorization": f"Bearer {token}"}
|
||||||
|
|
||||||
|
|
||||||
|
def _model_payload(prompt: float, provider_id: int, enabled: bool = True) -> dict:
|
||||||
|
return {
|
||||||
|
"id": "propagation-test-model",
|
||||||
|
"name": "Propagation Test Model",
|
||||||
|
"description": "model used to verify price propagation",
|
||||||
|
"created": 0,
|
||||||
|
"context_length": 128000,
|
||||||
|
"architecture": {
|
||||||
|
"modality": "text",
|
||||||
|
"input_modalities": ["text"],
|
||||||
|
"output_modalities": ["text"],
|
||||||
|
"tokenizer": "unknown",
|
||||||
|
"instruct_type": None,
|
||||||
|
},
|
||||||
|
"pricing": {
|
||||||
|
"prompt": prompt,
|
||||||
|
"completion": prompt * 2,
|
||||||
|
"input_cache_read": 0.0,
|
||||||
|
"input_cache_write": 0.0,
|
||||||
|
"request": 0.0,
|
||||||
|
"image": 0.0,
|
||||||
|
"web_search": 0.0,
|
||||||
|
"internal_reasoning": 0.0,
|
||||||
|
},
|
||||||
|
"per_request_limits": None,
|
||||||
|
"top_provider": None,
|
||||||
|
"upstream_provider_id": provider_id,
|
||||||
|
"canonical_slug": None,
|
||||||
|
"alias_ids": [],
|
||||||
|
"enabled": enabled,
|
||||||
|
"forwarded_model_id": "propagation-test-model",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_provider(session: AsyncSession, *, fee: float = 1.0) -> int:
|
||||||
|
"""Insert a provider, refresh the upstream map, and return its primary key."""
|
||||||
|
provider = UpstreamProviderRow(
|
||||||
|
provider_type="generic",
|
||||||
|
base_url="https://propagation-test.example/v1",
|
||||||
|
api_key="test-key",
|
||||||
|
provider_fee=fee,
|
||||||
|
)
|
||||||
|
session.add(provider)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(provider)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
assert provider.id is not None
|
||||||
|
return provider.id
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_price_edit_propagates_to_served_catalogue(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A price edit through the admin API must change the served /v1/models price."""
|
||||||
|
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
|
||||||
|
headers = _admin_headers()
|
||||||
|
r = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=headers,
|
||||||
|
json=_model_payload(prompt=1.0e-7, provider_id=provider_id),
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
|
||||||
|
# -- record the served price before edit -----------------------------------
|
||||||
|
public = await integration_client.get("/v1/models")
|
||||||
|
assert public.status_code == 200
|
||||||
|
public_data = public.json()
|
||||||
|
assert len(public_data["data"]) > 0, "catalogue must not be empty"
|
||||||
|
served_before = {
|
||||||
|
m["id"]: m.get("pricing", {}).get("prompt") for m in public_data["data"]
|
||||||
|
}
|
||||||
|
assert "propagation-test-model" in served_before
|
||||||
|
before = served_before["propagation-test-model"]
|
||||||
|
|
||||||
|
# -- edit the price and re-check -------------------------------------------
|
||||||
|
r = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=headers,
|
||||||
|
json=_model_payload(prompt=5.0e-7, provider_id=provider_id),
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
|
||||||
|
public = await integration_client.get("/v1/models")
|
||||||
|
assert public.status_code == 200
|
||||||
|
served_after = {
|
||||||
|
m["id"]: m.get("pricing", {}).get("prompt") for m in public.json()["data"]
|
||||||
|
}
|
||||||
|
after = served_after["propagation-test-model"]
|
||||||
|
|
||||||
|
assert before != after, "served price did not change after admin edit"
|
||||||
|
# With provider_fee=1.0 the served price equals the stored raw price.
|
||||||
|
assert after == pytest.approx(5.0e-7)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_admin_readback_is_raw_served_is_fee_adjusted(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""Admin read-back returns the raw price; /v1/models returns the fee-adjusted one."""
|
||||||
|
|
||||||
|
provider_id = await _seed_provider(integration_session, fee=1.05)
|
||||||
|
|
||||||
|
headers = _admin_headers()
|
||||||
|
model_payload = _model_payload(prompt=1.0e-7, provider_id=provider_id)
|
||||||
|
r = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=headers,
|
||||||
|
json=model_payload,
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
|
||||||
|
# Admin read-back: apply_provider_fee=False
|
||||||
|
admin_r = await integration_client.get(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models/propagation-test-model",
|
||||||
|
headers=headers,
|
||||||
|
)
|
||||||
|
assert admin_r.status_code == 200
|
||||||
|
admin_body = admin_r.json()
|
||||||
|
raw_prompt = admin_body["pricing"]["prompt"]
|
||||||
|
assert raw_prompt == pytest.approx(1.0e-7)
|
||||||
|
|
||||||
|
# Public /v1/models: fee-adjusted
|
||||||
|
public = await integration_client.get("/v1/models")
|
||||||
|
assert public.status_code == 200
|
||||||
|
served = {
|
||||||
|
m["id"]: m.get("pricing", {}).get("prompt") for m in public.json()["data"]
|
||||||
|
}
|
||||||
|
assert served["propagation-test-model"] == pytest.approx(1.0e-7 * 1.05)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_disabled_model_not_served(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A disabled model must be absent from /v1/models but still present in the DB."""
|
||||||
|
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
|
||||||
|
headers = _admin_headers()
|
||||||
|
r = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=headers,
|
||||||
|
json=_model_payload(prompt=1.0e-7, provider_id=provider_id, enabled=True),
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
|
||||||
|
# Confirm it appears in the public catalogue.
|
||||||
|
public = await integration_client.get("/v1/models")
|
||||||
|
served_ids = {m["id"] for m in public.json()["data"]}
|
||||||
|
assert "propagation-test-model" in served_ids
|
||||||
|
|
||||||
|
# -- disable via upsert ----------------------------------------------------
|
||||||
|
r = await integration_client.post(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=headers,
|
||||||
|
json=_model_payload(prompt=1.0e-7, provider_id=provider_id, enabled=False),
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
|
||||||
|
# Public catalogue must no longer list it.
|
||||||
|
public = await integration_client.get("/v1/models")
|
||||||
|
served_ids = {m["id"] for m in public.json()["data"]}
|
||||||
|
assert "propagation-test-model" not in served_ids
|
||||||
|
|
||||||
|
# DB row must still exist.
|
||||||
|
row = await integration_session.get(
|
||||||
|
ModelRow, ("propagation-test-model", provider_id)
|
||||||
|
)
|
||||||
|
assert row is not None
|
||||||
|
assert row.enabled is False
|
||||||
@@ -0,0 +1,435 @@
|
|||||||
|
"""Characterization tests for the model-serialisation pipeline ``_row_to_model`` runs.
|
||||||
|
|
||||||
|
These pin what the three public surfaces that reach it serve today: ``GET
|
||||||
|
/v1/models``, the admin single-model read-back, and the admin provider model
|
||||||
|
listing. Covered are the plain model, litellm cache backfill, preservation of
|
||||||
|
explicit cache rates, the request-price floor, the provider-fee flag against
|
||||||
|
recomputed max costs, survival of a failed sats conversion, the full serialised
|
||||||
|
dict field for field, and agreement between the two admin views.
|
||||||
|
|
||||||
|
Nothing here asserts on the deterministic USD half in isolation, so the pins hold
|
||||||
|
whether or not it is split out from the live BTC-rate conversion.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from collections.abc import Iterator
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from httpx import AsyncClient
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from routstr.core.admin import admin_sessions
|
||||||
|
from routstr.core.db import ModelRow, UpstreamProviderRow
|
||||||
|
from routstr.proxy import reinitialize_upstreams
|
||||||
|
|
||||||
|
|
||||||
|
# The conftest patches ``routstr.payment.price.sats_usd_price``, but
|
||||||
|
# ``models.py`` imports it as ``from .price import sats_usd_price`` — a
|
||||||
|
# local binding the conftest-level patch cannot reach. Pin it here so
|
||||||
|
# every test that goes through ``_row_to_model`` gets a real sats price.
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _pin_sats_usd() -> Iterator[None]:
|
||||||
|
with patch("routstr.payment.models.sats_usd_price", return_value=0.0005):
|
||||||
|
yield
|
||||||
|
|
||||||
|
|
||||||
|
def _admin_headers() -> dict[str, str]:
|
||||||
|
token = "test-serialisation-token"
|
||||||
|
admin_sessions[token] = int(
|
||||||
|
(datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp()
|
||||||
|
)
|
||||||
|
return {"Authorization": f"Bearer {token}"}
|
||||||
|
|
||||||
|
|
||||||
|
# -- helpers -------------------------------------------------------------------
|
||||||
|
|
||||||
|
_SEEDED_MODEL_ID = "ser-test-model"
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_provider(session: AsyncSession, fee: float = 1.0) -> int:
|
||||||
|
"""Insert a provider, refresh the upstream map, and return its primary key."""
|
||||||
|
provider = UpstreamProviderRow(
|
||||||
|
provider_type="generic",
|
||||||
|
base_url="https://serialisation-test.example/v1",
|
||||||
|
api_key="test-key",
|
||||||
|
provider_fee=fee,
|
||||||
|
)
|
||||||
|
session.add(provider)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(provider)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
assert provider.id is not None
|
||||||
|
return provider.id
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_model(
|
||||||
|
session: AsyncSession,
|
||||||
|
provider_id: int,
|
||||||
|
*,
|
||||||
|
model_id: str = _SEEDED_MODEL_ID,
|
||||||
|
prompt: float = 1.0e-7,
|
||||||
|
completion: float = 2.0e-7,
|
||||||
|
cache_read: float = 0.0,
|
||||||
|
cache_write: float = 0.0,
|
||||||
|
request_price: float = 0.0,
|
||||||
|
enabled: bool = True,
|
||||||
|
) -> ModelRow:
|
||||||
|
row = ModelRow(
|
||||||
|
id=model_id,
|
||||||
|
name=f"SerTest {model_id}",
|
||||||
|
description="characterization model",
|
||||||
|
created=0,
|
||||||
|
context_length=128000,
|
||||||
|
architecture=json.dumps(
|
||||||
|
{
|
||||||
|
"modality": "text",
|
||||||
|
"input_modalities": ["text"],
|
||||||
|
"output_modalities": ["text"],
|
||||||
|
"tokenizer": "unknown",
|
||||||
|
"instruct_type": None,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
pricing=json.dumps(
|
||||||
|
{
|
||||||
|
"prompt": prompt,
|
||||||
|
"completion": completion,
|
||||||
|
"input_cache_read": cache_read,
|
||||||
|
"input_cache_write": cache_write,
|
||||||
|
"request": request_price,
|
||||||
|
"image": 0.0,
|
||||||
|
"web_search": 0.0,
|
||||||
|
"internal_reasoning": 0.0,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
upstream_provider_id=provider_id,
|
||||||
|
enabled=enabled,
|
||||||
|
forwarded_model_id=model_id,
|
||||||
|
)
|
||||||
|
session.add(row)
|
||||||
|
await session.commit()
|
||||||
|
return row
|
||||||
|
|
||||||
|
|
||||||
|
async def _raw_via_admin(client: AsyncClient, provider_id: int, model_id: str) -> dict:
|
||||||
|
"""Return the raw (``apply_provider_fee=False``) model dict from admin read-back."""
|
||||||
|
r = await client.get(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models/{model_id}",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
return r.json()
|
||||||
|
|
||||||
|
|
||||||
|
async def _served_via_public(client: AsyncClient, model_id: str) -> dict | None:
|
||||||
|
"""Return the served model dict from /v1/models, or None if absent."""
|
||||||
|
r = await client.get("/v1/models")
|
||||||
|
assert r.status_code == 200
|
||||||
|
return {m["id"]: m for m in r.json()["data"]}.get(model_id)
|
||||||
|
|
||||||
|
|
||||||
|
# -- test 1: plain model -------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_plain_model_serialisation(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A model with prompt + completion prices has the expected serialised shape."""
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
await _seed_model(integration_session, provider_id)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
|
||||||
|
# Admin read-back (raw, no fee)
|
||||||
|
body = await _raw_via_admin(
|
||||||
|
integration_client,
|
||||||
|
provider_id,
|
||||||
|
_SEEDED_MODEL_ID,
|
||||||
|
)
|
||||||
|
assert body["id"] == _SEEDED_MODEL_ID
|
||||||
|
assert body["pricing"]["prompt"] == pytest.approx(1.0e-7)
|
||||||
|
assert body["pricing"]["completion"] == pytest.approx(2.0e-7)
|
||||||
|
assert body["sats_pricing"] is not None
|
||||||
|
# With sats_usd_price = 0.0005 (the fixture)
|
||||||
|
assert body["sats_pricing"]["prompt"] == pytest.approx(1.0e-7 / 0.0005)
|
||||||
|
assert body["sats_pricing"]["completion"] == pytest.approx(2.0e-7 / 0.0005)
|
||||||
|
|
||||||
|
# Public /v1/models (fee applied)
|
||||||
|
s = await _served_via_public(integration_client, _SEEDED_MODEL_ID)
|
||||||
|
assert s is not None, f"{_SEEDED_MODEL_ID} not found in /v1/models"
|
||||||
|
# fee=1.0 so values match raw
|
||||||
|
assert s["pricing"]["prompt"] == pytest.approx(1.0e-7)
|
||||||
|
assert s["pricing"]["completion"] == pytest.approx(2.0e-7)
|
||||||
|
|
||||||
|
|
||||||
|
# -- test 2: no cache rates (litellm backfill) ---------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_model_without_cache_rates_gets_litellm_backfill(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A well-known model without cache rates gets them from litellm's cost map."""
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
# Use a real litellm-known id so backfill_cache_pricing can find it.
|
||||||
|
await _seed_model(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="gpt-4o",
|
||||||
|
prompt=2.5e-6,
|
||||||
|
completion=1.0e-5,
|
||||||
|
cache_read=0.0,
|
||||||
|
cache_write=0.0,
|
||||||
|
)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
|
||||||
|
# Admin read-back (raw): cache_read should be present after backfill.
|
||||||
|
# (cache_write may not be in litellm's map for every model.)
|
||||||
|
body = await _raw_via_admin(
|
||||||
|
integration_client,
|
||||||
|
provider_id,
|
||||||
|
"gpt-4o",
|
||||||
|
)
|
||||||
|
assert body["pricing"]["input_cache_read"] > 0.0, (
|
||||||
|
"backfill_cache_pricing should have filled input_cache_read from litellm"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# -- test 3: cache rates already present (not overwritten) ---------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_existing_cache_rates_not_overwritten(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A model with explicit cache rates must keep them; the backfill is a no-op."""
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
await _seed_model(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
prompt=1.0e-7,
|
||||||
|
completion=2.0e-7,
|
||||||
|
cache_read=9.99e-9,
|
||||||
|
cache_write=8.88e-9,
|
||||||
|
)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
|
||||||
|
body = await _raw_via_admin(
|
||||||
|
integration_client,
|
||||||
|
provider_id,
|
||||||
|
_SEEDED_MODEL_ID,
|
||||||
|
)
|
||||||
|
assert body["pricing"]["input_cache_read"] == pytest.approx(9.99e-9)
|
||||||
|
assert body["pricing"]["input_cache_write"] == pytest.approx(8.88e-9)
|
||||||
|
|
||||||
|
|
||||||
|
# -- test 4: request price floor -----------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_model_with_request_price(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A model with a request price floor carries it through to the served model."""
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
await _seed_model(integration_session, provider_id, request_price=0.01)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
|
||||||
|
body = await _raw_via_admin(
|
||||||
|
integration_client,
|
||||||
|
provider_id,
|
||||||
|
_SEEDED_MODEL_ID,
|
||||||
|
)
|
||||||
|
assert body["pricing"]["request"] == pytest.approx(0.01)
|
||||||
|
|
||||||
|
|
||||||
|
# -- test 5: provider_fee=True vs False, max costs recomputed ------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_fee_flag_changes_pricing_but_max_costs_are_recomputed(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""With provider_fee=1.5, fee-adjusted pricing is 1.5× raw, but max costs are NOT."""
|
||||||
|
provider_id = await _seed_provider(integration_session, fee=1.5)
|
||||||
|
await _seed_model(integration_session, provider_id)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
|
||||||
|
# Admin read-back: raw, no fee.
|
||||||
|
body = await _raw_via_admin(
|
||||||
|
integration_client,
|
||||||
|
provider_id,
|
||||||
|
_SEEDED_MODEL_ID,
|
||||||
|
)
|
||||||
|
assert body["pricing"]["prompt"] == pytest.approx(1.0e-7)
|
||||||
|
|
||||||
|
# Public /v1/models: fee applied.
|
||||||
|
s = await _served_via_public(integration_client, _SEEDED_MODEL_ID)
|
||||||
|
assert s is not None
|
||||||
|
assert s["pricing"]["prompt"] == pytest.approx(1.0e-7 * 1.5)
|
||||||
|
|
||||||
|
# max_prompt_cost must NOT just be multiplied by 1.5 — it is recomputed from
|
||||||
|
# the fee-inflated per-token rates and context_length.
|
||||||
|
cl = 128_000
|
||||||
|
expected_max_prompt = cl * 1.0e-7 * 1.5
|
||||||
|
assert s["pricing"]["max_prompt_cost"] == pytest.approx(expected_max_prompt)
|
||||||
|
|
||||||
|
|
||||||
|
# -- test 6: sats conversion failure keeps the model alive ---------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_model_survives_sats_conversion_failure(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""When the BTC feed fails, the model still returns — with no sats_pricing."""
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
await _seed_model(integration_session, provider_id)
|
||||||
|
|
||||||
|
# Make sats_usd_price raise so _update_model_sats_pricing swallows it.
|
||||||
|
# The admin read-back must happen inside the patch block.
|
||||||
|
with patch(
|
||||||
|
"routstr.payment.models.sats_usd_price",
|
||||||
|
side_effect=RuntimeError("BTC feed down"),
|
||||||
|
):
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
body = await _raw_via_admin(
|
||||||
|
integration_client,
|
||||||
|
provider_id,
|
||||||
|
_SEEDED_MODEL_ID,
|
||||||
|
)
|
||||||
|
assert body["id"] == _SEEDED_MODEL_ID
|
||||||
|
assert body["sats_pricing"] is None, (
|
||||||
|
"sats conversion failure must not crash — the model returns with no sats_pricing"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# -- test 7: the whole serialised dict -----------------------------------------
|
||||||
|
|
||||||
|
# The sats figures are the USD ones divided by the pinned 0.0005 rate; written
|
||||||
|
# as the division so the expectation carries the same float error the code does.
|
||||||
|
_SATS = 0.0005
|
||||||
|
|
||||||
|
|
||||||
|
def _expected_serialised_model(provider_id: int) -> dict:
|
||||||
|
"""Every field the raw admin read-back produces for the seeded model.
|
||||||
|
|
||||||
|
Note the max-cost asymmetry: with fees off the USD max costs stay at zero
|
||||||
|
while their sats counterparts are computed. That is what the code does
|
||||||
|
today, and pinning it is the point.
|
||||||
|
"""
|
||||||
|
return {
|
||||||
|
"alias_ids": None,
|
||||||
|
"architecture": {
|
||||||
|
"input_modalities": ["text"],
|
||||||
|
"instruct_type": None,
|
||||||
|
"modality": "text",
|
||||||
|
"output_modalities": ["text"],
|
||||||
|
"tokenizer": "unknown",
|
||||||
|
},
|
||||||
|
"canonical_slug": None,
|
||||||
|
"context_length": 128000,
|
||||||
|
"created": 0,
|
||||||
|
"description": "characterization model",
|
||||||
|
"enabled": True,
|
||||||
|
"forwarded_model_id": _SEEDED_MODEL_ID,
|
||||||
|
"id": _SEEDED_MODEL_ID,
|
||||||
|
"name": f"SerTest {_SEEDED_MODEL_ID}",
|
||||||
|
"per_request_limits": None,
|
||||||
|
"pricing": {
|
||||||
|
"completion": 2.0e-7,
|
||||||
|
"image": 0.0,
|
||||||
|
"input_cache_read": 0.0,
|
||||||
|
"input_cache_write": 0.0,
|
||||||
|
"internal_reasoning": 0.0,
|
||||||
|
"max_completion_cost": 0.0,
|
||||||
|
"max_cost": 0.0,
|
||||||
|
"max_prompt_cost": 0.0,
|
||||||
|
"prompt": 1.0e-7,
|
||||||
|
"request": 0.01,
|
||||||
|
"web_search": 0.0,
|
||||||
|
},
|
||||||
|
"sats_pricing": {
|
||||||
|
"completion": 2.0e-7 / _SATS,
|
||||||
|
"image": 0.0,
|
||||||
|
"input_cache_read": 0.0,
|
||||||
|
"input_cache_write": 0.0,
|
||||||
|
"internal_reasoning": 0.0,
|
||||||
|
"max_completion_cost": 0.0,
|
||||||
|
"max_cost": 0.001,
|
||||||
|
"max_prompt_cost": 0.0,
|
||||||
|
"prompt": 1.0e-7 / _SATS,
|
||||||
|
"request": 0.01 / _SATS,
|
||||||
|
"web_search": 0.0,
|
||||||
|
},
|
||||||
|
"top_provider": None,
|
||||||
|
"upstream_provider_id": provider_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_full_serialised_model_is_unchanged(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""Pin every field of the serialised model, not just the interesting ones.
|
||||||
|
|
||||||
|
The tests above pin values. A refactor that dropped a field outright
|
||||||
|
would satisfy all of them and fail only here.
|
||||||
|
"""
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
await _seed_model(integration_session, provider_id, request_price=0.01)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
|
||||||
|
body = await _raw_via_admin(
|
||||||
|
integration_client,
|
||||||
|
provider_id,
|
||||||
|
_SEEDED_MODEL_ID,
|
||||||
|
)
|
||||||
|
assert body == _expected_serialised_model(provider_id)
|
||||||
|
|
||||||
|
|
||||||
|
# -- test 8: the provider model listing ----------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_provider_listing_matches_single_model_read_back(
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""The listing is a third entry into the same builder — it must agree."""
|
||||||
|
provider_id = await _seed_provider(integration_session)
|
||||||
|
await _seed_model(integration_session, provider_id, request_price=0.01)
|
||||||
|
await reinitialize_upstreams()
|
||||||
|
|
||||||
|
single = await _raw_via_admin(
|
||||||
|
integration_client,
|
||||||
|
provider_id,
|
||||||
|
_SEEDED_MODEL_ID,
|
||||||
|
)
|
||||||
|
|
||||||
|
r = await integration_client.get(
|
||||||
|
f"/admin/api/upstream-providers/{provider_id}/models",
|
||||||
|
headers=_admin_headers(),
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
listed = {m["id"]: m for m in r.json()["db_models"]}
|
||||||
|
assert _SEEDED_MODEL_ID in listed, "seeded model missing from the provider listing"
|
||||||
|
assert listed[_SEEDED_MODEL_ID] == single
|
||||||
@@ -0,0 +1,747 @@
|
|||||||
|
"""Regression tests for the production "negative available balance" 402.
|
||||||
|
|
||||||
|
Invariants protected:
|
||||||
|
* billing and the admin API agree on what "available" means
|
||||||
|
* an in-flight (heartbeaten) reservation is never swept; an abandoned one is
|
||||||
|
released and can never be charged afterwards
|
||||||
|
* a cost overrun spends only its own reservation plus unreserved balance
|
||||||
|
* corrupt reservations are repaired terminally instead of poisoning cleanup
|
||||||
|
* under concurrency, balance / reserved / available never go negative
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import random
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from typing import Awaitable, Callable
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from httpx import AsyncClient
|
||||||
|
from sqlmodel import col, func, select, update
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from routstr.core.db import ApiKey, ReservationRelease
|
||||||
|
from routstr.payment.cost_calculation import CostData
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.integration
|
||||||
|
|
||||||
|
# Realistic sweeper timeout: a renewed (heartbeaten) reservation stays alive,
|
||||||
|
# a reservation backdated past this is released.
|
||||||
|
STALE_TIMEOUT_SECONDS = 300
|
||||||
|
|
||||||
|
# created_at is whole seconds; -1 makes every reservation stale immediately.
|
||||||
|
# Only used in the fuzz test to stress the terminal-release path.
|
||||||
|
SWEEP_EVERYTHING = -1
|
||||||
|
|
||||||
|
|
||||||
|
def _cost_data(total_msats: int) -> CostData:
|
||||||
|
return CostData(
|
||||||
|
base_msats=0,
|
||||||
|
input_msats=total_msats // 2,
|
||||||
|
output_msats=total_msats - total_msats // 2,
|
||||||
|
total_msats=total_msats,
|
||||||
|
total_usd=0.0,
|
||||||
|
input_tokens=100,
|
||||||
|
output_tokens=100,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _response(model: str = "test-model") -> dict:
|
||||||
|
return {
|
||||||
|
"model": model,
|
||||||
|
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _new_key(session: AsyncSession, balance: int) -> str:
|
||||||
|
key_hash = f"test_neg_{uuid.uuid4().hex}"
|
||||||
|
session.add(
|
||||||
|
ApiKey(
|
||||||
|
hashed_key=key_hash,
|
||||||
|
balance=balance,
|
||||||
|
reserved_balance=0,
|
||||||
|
total_spent=0,
|
||||||
|
total_requests=0,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return key_hash
|
||||||
|
|
||||||
|
|
||||||
|
async def _backdate_reservation(
|
||||||
|
session: AsyncSession, release_id: str, seconds: int
|
||||||
|
) -> None:
|
||||||
|
"""Age a reservation's lease as if it had not been renewed for `seconds`."""
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ReservationRelease)
|
||||||
|
.where(col(ReservationRelease.id) == release_id)
|
||||||
|
.values(created_at=col(ReservationRelease.created_at) - seconds)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def _wait_for(
|
||||||
|
predicate: "Callable[[], Awaitable[bool]]",
|
||||||
|
timeout: float = 10.0,
|
||||||
|
interval: float = 0.1,
|
||||||
|
) -> bool:
|
||||||
|
"""Bounded polling instead of fixed sleeps for background-task effects."""
|
||||||
|
deadline = asyncio.get_event_loop().time() + timeout
|
||||||
|
while True:
|
||||||
|
if await predicate():
|
||||||
|
return True
|
||||||
|
if asyncio.get_event_loop().time() > deadline:
|
||||||
|
return False
|
||||||
|
await asyncio.sleep(interval)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_402_reports_negative_available_while_admin_shows_positive(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
integration_client: AsyncClient,
|
||||||
|
) -> None:
|
||||||
|
"""The production symptom: billing rejects on balance - reserved_balance,
|
||||||
|
so the admin endpoint must expose reserved/available, not just balance."""
|
||||||
|
from fastapi import HTTPException
|
||||||
|
|
||||||
|
from routstr.auth import _validate_bearer_key_locked
|
||||||
|
from routstr.core.db import set_admin_password
|
||||||
|
|
||||||
|
key_hash = await _new_key(integration_session, balance=263_000)
|
||||||
|
# Leaked reservations slightly exceeding the balance, as in production.
|
||||||
|
await integration_session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.values(reserved_balance=267_215)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
with pytest.raises(HTTPException) as exc:
|
||||||
|
await _validate_bearer_key_locked(
|
||||||
|
"sk-" + key_hash, integration_session, min_cost=1
|
||||||
|
)
|
||||||
|
|
||||||
|
assert exc.value.status_code == 402
|
||||||
|
message = exc.value.detail["error"]["message"] # type: ignore[index]
|
||||||
|
assert "-4.215 sats (-4215 msats) available" in message, message
|
||||||
|
|
||||||
|
await set_admin_password(integration_session, "test-admin-pw")
|
||||||
|
login = await integration_client.post(
|
||||||
|
"/admin/api/login", json={"password": "test-admin-pw"}
|
||||||
|
)
|
||||||
|
assert login.status_code == 200
|
||||||
|
token = login.json()["token"]
|
||||||
|
resp = await integration_client.get(
|
||||||
|
"/admin/api/temporary-balances",
|
||||||
|
params={"search": key_hash},
|
||||||
|
headers={"Authorization": f"Bearer {token}"},
|
||||||
|
)
|
||||||
|
assert resp.status_code == 200
|
||||||
|
rows = [row for row in resp.json()["balances"] if row["hashed_key"] == key_hash]
|
||||||
|
assert len(rows) == 1
|
||||||
|
row = rows[0]
|
||||||
|
assert row["balance"] == 263_000
|
||||||
|
assert row["reserved_balance"] == 267_215
|
||||||
|
assert row["available_balance"] == -4_215
|
||||||
|
assert resp.json()["totals"] == {
|
||||||
|
"total_balance": 263_000,
|
||||||
|
"total_reserved_balance": 267_215,
|
||||||
|
"total_available_balance": -4_215,
|
||||||
|
"total_spent": 0,
|
||||||
|
"total_requests": 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_abandoned_reservation_is_swept_and_cannot_finalize(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""Release is terminal: a swept reservation's late finalizer must not
|
||||||
|
charge, or it could spend funds since reserved by another request."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import release_stale_reservations
|
||||||
|
|
||||||
|
cost = 5_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
# Client vanished: the lease is never renewed and ages past the timeout.
|
||||||
|
await _backdate_reservation(
|
||||||
|
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 1
|
||||||
|
)
|
||||||
|
released = await release_stale_reservations(
|
||||||
|
integration_session, STALE_TIMEOUT_SECONDS
|
||||||
|
)
|
||||||
|
assert released == 1
|
||||||
|
|
||||||
|
# A zombie finalizer shows up afterwards; it must not charge.
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
record = await integration_session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
assert key.total_spent == 0
|
||||||
|
assert key.balance == 10_000
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_sweeper_cannot_release_a_reservation_renewed_after_selection(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A heartbeat landing between the sweeper's select and its transition
|
||||||
|
must win — exercised through the public sweeper entry point."""
|
||||||
|
from routstr.auth import (
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
renew_reservation,
|
||||||
|
)
|
||||||
|
from routstr.core import db as core_db
|
||||||
|
from routstr.core.db import release_stale_reservations
|
||||||
|
|
||||||
|
key_hash = await _new_key(integration_session, balance=1_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, 1_000, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
await _backdate_reservation(
|
||||||
|
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 100
|
||||||
|
)
|
||||||
|
|
||||||
|
real_transition = core_db._transition_stale_reservation
|
||||||
|
|
||||||
|
async def renew_then_transition(
|
||||||
|
session: AsyncSession, reservation_id: str, cutoff: int
|
||||||
|
) -> bool:
|
||||||
|
# The sweeper selected this reservation as stale; the heartbeat
|
||||||
|
# renews exactly between that select and the transition.
|
||||||
|
assert await renew_reservation(reservation, session)
|
||||||
|
return await real_transition(session, reservation_id, cutoff)
|
||||||
|
|
||||||
|
with patch.object(core_db, "_transition_stale_reservation", renew_then_transition):
|
||||||
|
released = await release_stale_reservations(
|
||||||
|
integration_session, STALE_TIMEOUT_SECONDS
|
||||||
|
)
|
||||||
|
|
||||||
|
assert released == 0
|
||||||
|
record = await integration_session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "active"
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.reserved_balance == 1_000
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_legacy_cleanup_cannot_erase_a_reservation_committed_after_its_read(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A reservation committing between the legacy sweep's read and its
|
||||||
|
zeroing must survive — exercised through the public sweeper entry point."""
|
||||||
|
from routstr.core import db as core_db
|
||||||
|
from routstr.core.db import release_stale_reservations
|
||||||
|
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
stale_reserved_at = int(time.time()) - (STALE_TIMEOUT_SECONDS + 100)
|
||||||
|
await integration_session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.values(reserved_balance=500, reserved_at=stale_reserved_at)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
real_release = core_db._release_legacy_aggregate
|
||||||
|
|
||||||
|
async def commit_reservation_then_release(
|
||||||
|
session: AsyncSession,
|
||||||
|
target_key_hash: str,
|
||||||
|
observed_reserved: int,
|
||||||
|
observed_reserved_at: int | None,
|
||||||
|
) -> bool:
|
||||||
|
# The sweeper read the legacy aggregate and found no active durable
|
||||||
|
# owner; a new reservation commits exactly before the zeroing lands.
|
||||||
|
session.add(
|
||||||
|
ReservationRelease(
|
||||||
|
id=uuid.uuid4().hex,
|
||||||
|
key_hash=target_key_hash,
|
||||||
|
billing_key_hash=target_key_hash,
|
||||||
|
reserved_msats=700,
|
||||||
|
status="active",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == target_key_hash)
|
||||||
|
.values(
|
||||||
|
reserved_balance=col(ApiKey.reserved_balance) + 700,
|
||||||
|
reserved_at=int(time.time()),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return await real_release(
|
||||||
|
session, target_key_hash, observed_reserved, observed_reserved_at
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
core_db, "_release_legacy_aggregate", commit_reservation_then_release
|
||||||
|
):
|
||||||
|
released = await release_stale_reservations(
|
||||||
|
integration_session, STALE_TIMEOUT_SECONDS
|
||||||
|
)
|
||||||
|
|
||||||
|
assert released == 0
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.reserved_balance == 1_200, "legacy cleanup erased a live reservation"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_sweeper_repairs_corrupt_reservation_and_continues_batch(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""One corrupt durable reservation (aggregate no longer holds its msats)
|
||||||
|
must be terminalized without aggregate subtraction, and must not stop the
|
||||||
|
rest of the batch from being released normally."""
|
||||||
|
from routstr.auth import get_reservation_snapshot, pay_for_request
|
||||||
|
from routstr.core.db import release_stale_reservations
|
||||||
|
|
||||||
|
cost = 1_000
|
||||||
|
corrupt_hash = await _new_key(integration_session, balance=cost)
|
||||||
|
healthy_hash = await _new_key(integration_session, balance=cost)
|
||||||
|
|
||||||
|
corrupt_key = await integration_session.get(ApiKey, corrupt_hash)
|
||||||
|
assert corrupt_key is not None
|
||||||
|
await pay_for_request(corrupt_key, cost, integration_session)
|
||||||
|
corrupt_reservation = await get_reservation_snapshot(
|
||||||
|
corrupt_key, integration_session
|
||||||
|
)
|
||||||
|
|
||||||
|
healthy_key = await integration_session.get(ApiKey, healthy_hash)
|
||||||
|
assert healthy_key is not None
|
||||||
|
await pay_for_request(healthy_key, cost, integration_session)
|
||||||
|
healthy_reservation = await get_reservation_snapshot(
|
||||||
|
healthy_key, integration_session
|
||||||
|
)
|
||||||
|
|
||||||
|
# Corrupt the first key: its aggregate no longer holds the reservation.
|
||||||
|
await integration_session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == corrupt_hash)
|
||||||
|
.values(reserved_balance=0)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
for reservation in (corrupt_reservation, healthy_reservation):
|
||||||
|
await _backdate_reservation(
|
||||||
|
integration_session, reservation.release_id, STALE_TIMEOUT_SECONDS + 100
|
||||||
|
)
|
||||||
|
|
||||||
|
released = await release_stale_reservations(
|
||||||
|
integration_session, STALE_TIMEOUT_SECONDS
|
||||||
|
)
|
||||||
|
assert released == 2, "a corrupt record must not abort the sweep batch"
|
||||||
|
|
||||||
|
for reservation in (corrupt_reservation, healthy_reservation):
|
||||||
|
record = await integration_session.get(
|
||||||
|
ReservationRelease, reservation.release_id
|
||||||
|
)
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
healthy_key = await integration_session.get(ApiKey, healthy_hash)
|
||||||
|
corrupt_key = await integration_session.get(ApiKey, corrupt_hash)
|
||||||
|
assert healthy_key is not None and corrupt_key is not None
|
||||||
|
assert healthy_key.reserved_balance == 0
|
||||||
|
assert corrupt_key.reserved_balance == 0
|
||||||
|
assert corrupt_key.balance == cost, "repair must not touch balances"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_heartbeat_survives_a_rolled_back_charge_attempt(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
"""Claiming a reservation must not stop its heartbeat: a rollback restores
|
||||||
|
the active reservation, which then still needs lease renewal."""
|
||||||
|
from routstr import auth
|
||||||
|
from routstr.auth import (
|
||||||
|
_claim_reservation_for_charge,
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import create_session
|
||||||
|
|
||||||
|
cost = 1_000
|
||||||
|
timeout = 3 # heartbeat interval = 1s
|
||||||
|
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=2_000)
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, cost, session)
|
||||||
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# A charge attempt claims the reservation, then its transaction
|
||||||
|
# fails and rolls back.
|
||||||
|
async with create_session() as session:
|
||||||
|
assert await _claim_reservation_for_charge(reservation, session)
|
||||||
|
await session.rollback()
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "active"
|
||||||
|
await _backdate_reservation(
|
||||||
|
session, reservation.release_id, timeout * 10
|
||||||
|
)
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None
|
||||||
|
backdated_lease = record.created_at
|
||||||
|
|
||||||
|
async def lease_renewed() -> bool:
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(
|
||||||
|
ReservationRelease, reservation.release_id
|
||||||
|
)
|
||||||
|
return record is not None and record.created_at > backdated_lease
|
||||||
|
|
||||||
|
assert await _wait_for(lease_renewed), (
|
||||||
|
"heartbeat did not survive the rolled-back charge attempt"
|
||||||
|
)
|
||||||
|
|
||||||
|
# The restored reservation still finalizes normally.
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
await auth._stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == cost
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "charged"
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.total_spent == cost
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_heartbeat_dies_with_its_request_so_sweeper_can_recover(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
"""A request that vanishes without finalizing must not renew forever —
|
||||||
|
its heartbeat stops with the owning task and the sweeper reclaims the
|
||||||
|
funds."""
|
||||||
|
from routstr import auth
|
||||||
|
from routstr.auth import get_reservation_snapshot, pay_for_request
|
||||||
|
from routstr.core.db import create_session, release_stale_reservations
|
||||||
|
|
||||||
|
cost = 1_000
|
||||||
|
timeout = 3 # heartbeat interval = 1s
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=cost)
|
||||||
|
|
||||||
|
holder: dict = {}
|
||||||
|
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
|
||||||
|
|
||||||
|
async def doomed_request() -> None:
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, cost, session)
|
||||||
|
holder["reservation"] = await get_reservation_snapshot(key, session)
|
||||||
|
# ...request control dies here, no finalize and no release.
|
||||||
|
|
||||||
|
await asyncio.create_task(doomed_request())
|
||||||
|
release_id = holder["reservation"].release_id
|
||||||
|
|
||||||
|
try:
|
||||||
|
# While the heartbeat is still winding down it may renew once
|
||||||
|
# more; keep backdating until the sweeper wins, which it must as
|
||||||
|
# soon as the dead owner is noticed.
|
||||||
|
async def sweeper_recovered() -> bool:
|
||||||
|
async with create_session() as session:
|
||||||
|
await _backdate_reservation(session, release_id, timeout * 10)
|
||||||
|
return await release_stale_reservations(session, timeout) == 1
|
||||||
|
|
||||||
|
assert await _wait_for(sweeper_recovered), (
|
||||||
|
"sweeper never recovered the abandoned reservation"
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
await auth._stop_reservation_heartbeat(release_id)
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(ReservationRelease, release_id)
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
assert key.total_spent == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reservation_heartbeat_covers_the_whole_request_lifecycle(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
"""pay_for_request starts the heartbeat, finalization stops it, and a
|
||||||
|
backdated lease is renewed in the background without any manual call."""
|
||||||
|
from routstr import auth
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import create_session, release_stale_reservations
|
||||||
|
|
||||||
|
cost = 1_000
|
||||||
|
timeout = 3 # heartbeat interval = 1s
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=2 * cost)
|
||||||
|
|
||||||
|
with patch.object(auth.settings, "stale_reservation_timeout_seconds", timeout):
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, cost, session)
|
||||||
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Only a background renewal can keep this alive now.
|
||||||
|
async with create_session() as session:
|
||||||
|
await _backdate_reservation(
|
||||||
|
session, reservation.release_id, timeout * 10
|
||||||
|
)
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None
|
||||||
|
backdated_lease = record.created_at
|
||||||
|
|
||||||
|
async def lease_renewed() -> bool:
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(
|
||||||
|
ReservationRelease, reservation.release_id
|
||||||
|
)
|
||||||
|
return record is not None and record.created_at > backdated_lease
|
||||||
|
|
||||||
|
assert await _wait_for(lease_renewed), "heartbeat never renewed the lease"
|
||||||
|
async with create_session() as session:
|
||||||
|
assert await release_stale_reservations(session, timeout) == 0
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
await auth._stop_reservation_heartbeat(reservation.release_id)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == cost
|
||||||
|
async with create_session() as session:
|
||||||
|
record = await session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert record is not None and record.status == "charged"
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_overrun_cannot_spend_a_concurrent_reservation(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""An overrun is capped to its own reservation plus unreserved balance;
|
||||||
|
the sibling's reserved funds stay untouched and available stays >= 0."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved_each = 100
|
||||||
|
overrun_cost = 150 # A's real token cost exceeds its reservation
|
||||||
|
|
||||||
|
# Balance covers exactly two reservations; nothing free on top.
|
||||||
|
key_hash = await _new_key(integration_session, balance=2 * reserved_each)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved_each, integration_session)
|
||||||
|
reservation_a = await get_reservation_snapshot(key, integration_session)
|
||||||
|
await pay_for_request(key, reserved_each, integration_session)
|
||||||
|
await get_reservation_snapshot(key, integration_session) # B stays in flight
|
||||||
|
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.reserved_balance == 2 * reserved_each
|
||||||
|
|
||||||
|
# Only A finalizes; B is still streaming and its funds must stay reserved.
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(overrun_cost)):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved_each,
|
||||||
|
reservation_snapshot=reservation_a,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == reserved_each
|
||||||
|
assert result["total_msats"] == overrun_cost
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.total_spent == reserved_each
|
||||||
|
assert key.reserved_balance == reserved_each
|
||||||
|
assert key.balance == reserved_each
|
||||||
|
assert key.total_balance >= 0, (
|
||||||
|
f"available balance went negative: balance={key.balance} "
|
||||||
|
f"reserved={key.reserved_balance} -> {key.total_balance} msats; "
|
||||||
|
"the overrun charge consumed the still-reserved funds of request B"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_concurrent_requests_with_sweeper_keep_balance_invariants(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
"""Fuzz: concurrent requests against an everything-is-stale sweeper.
|
||||||
|
Some requests legitimately finish uncharged (release is terminal), but
|
||||||
|
balances must never go negative and no reservation may stay active."""
|
||||||
|
from fastapi import HTTPException
|
||||||
|
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import create_session, release_stale_reservations
|
||||||
|
|
||||||
|
rng = random.Random(1337)
|
||||||
|
starting_balance = 200_000
|
||||||
|
n_requests = 24
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=starting_balance)
|
||||||
|
|
||||||
|
completed_costs: list[int] = []
|
||||||
|
rejected_requests = 0
|
||||||
|
|
||||||
|
async def one_request(index: int) -> None:
|
||||||
|
nonlocal rejected_requests
|
||||||
|
reserved = rng.randrange(1_000, 4_000)
|
||||||
|
actual = max(1, int(reserved * rng.uniform(0.5, 1.1)))
|
||||||
|
try:
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, reserved, session)
|
||||||
|
except HTTPException as exc:
|
||||||
|
# A depleted balance is the only legitimate rejection.
|
||||||
|
assert exc.status_code == 402, exc.detail
|
||||||
|
rejected_requests += 1
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
except RuntimeError:
|
||||||
|
# The everything-is-stale sweeper can release the reservation
|
||||||
|
# before the stream even starts; the request aborts uncharged.
|
||||||
|
return
|
||||||
|
|
||||||
|
await asyncio.sleep(rng.uniform(0, 0.02)) # the "stream"
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(str(actual)),
|
||||||
|
session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
completed_costs.append(actual)
|
||||||
|
|
||||||
|
sweeping = True
|
||||||
|
|
||||||
|
async def sweeper() -> None:
|
||||||
|
while sweeping:
|
||||||
|
async with create_session() as session:
|
||||||
|
await release_stale_reservations(session, SWEEP_EVERYTHING)
|
||||||
|
await asyncio.sleep(0.002)
|
||||||
|
|
||||||
|
sweep_task = asyncio.create_task(sweeper())
|
||||||
|
try:
|
||||||
|
with patch(
|
||||||
|
"routstr.auth.calculate_cost",
|
||||||
|
side_effect=lambda response_data, *a, **k: _cost_data(
|
||||||
|
int(response_data["model"])
|
||||||
|
),
|
||||||
|
):
|
||||||
|
await asyncio.gather(*(one_request(i) for i in range(n_requests)))
|
||||||
|
finally:
|
||||||
|
sweeping = False
|
||||||
|
await sweep_task
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
leftover_active = (
|
||||||
|
await session.exec( # type: ignore[call-overload]
|
||||||
|
select(func.count())
|
||||||
|
.select_from(ReservationRelease)
|
||||||
|
.where(col(ReservationRelease.status) == "active")
|
||||||
|
)
|
||||||
|
).one()
|
||||||
|
|
||||||
|
assert completed_costs or rejected_requests, "no request made any progress"
|
||||||
|
assert key.balance >= 0, f"balance went negative: {key.balance}"
|
||||||
|
assert key.reserved_balance >= 0, (
|
||||||
|
f"reserved_balance went negative: {key.reserved_balance}"
|
||||||
|
)
|
||||||
|
assert key.total_balance >= 0, (
|
||||||
|
f"available balance negative: balance={key.balance} "
|
||||||
|
f"reserved={key.reserved_balance} (this is the production symptom)"
|
||||||
|
)
|
||||||
|
assert key.total_spent <= starting_balance, (
|
||||||
|
f"spent {key.total_spent} of a {starting_balance} balance"
|
||||||
|
)
|
||||||
|
# Requests whose reservation was swept mid-flight finish uncharged, so the
|
||||||
|
# charged total can only be at most the sum of completed request costs.
|
||||||
|
max_expected_spend = sum(completed_costs)
|
||||||
|
assert key.total_spent <= max_expected_spend, (
|
||||||
|
f"charged {key.total_spent} msats but completed requests only cost "
|
||||||
|
f"{max_expected_spend}"
|
||||||
|
)
|
||||||
|
assert leftover_active == 0, (
|
||||||
|
f"{leftover_active} reservations still active after all requests finished"
|
||||||
|
)
|
||||||
@@ -0,0 +1,450 @@
|
|||||||
|
"""Invariant coverage for the reserve → charge → release money path.
|
||||||
|
|
||||||
|
Every finalization branch of ``adjust_payment_for_tokens`` must respect the
|
||||||
|
same accounting rules: a completed request is charged exactly once, its
|
||||||
|
reported ``charged_msats`` matches the actual debit, it never spends more than
|
||||||
|
its own reservation leaves available, and concurrent requests cannot raid each
|
||||||
|
other's reservations.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import uuid
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sqlmodel import col, select, update
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from routstr.core.db import ApiKey, ReservationRelease
|
||||||
|
from routstr.payment.cost_calculation import (
|
||||||
|
CostData,
|
||||||
|
CostDataError,
|
||||||
|
MaxCostData,
|
||||||
|
)
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.integration
|
||||||
|
|
||||||
|
|
||||||
|
def _cost_data(total_msats: int, cls: type[CostData] = CostData) -> CostData:
|
||||||
|
return cls(
|
||||||
|
base_msats=0,
|
||||||
|
input_msats=total_msats // 2,
|
||||||
|
output_msats=total_msats - total_msats // 2,
|
||||||
|
total_msats=total_msats,
|
||||||
|
total_usd=0.0,
|
||||||
|
input_tokens=100,
|
||||||
|
output_tokens=100,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _response() -> dict:
|
||||||
|
return {
|
||||||
|
"model": "test-model",
|
||||||
|
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _new_key(session: AsyncSession, balance: int) -> str:
|
||||||
|
key_hash = f"test_inv_{uuid.uuid4().hex}"
|
||||||
|
session.add(
|
||||||
|
ApiKey(
|
||||||
|
hashed_key=key_hash,
|
||||||
|
balance=balance,
|
||||||
|
reserved_balance=0,
|
||||||
|
total_spent=0,
|
||||||
|
total_requests=0,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
return key_hash
|
||||||
|
|
||||||
|
|
||||||
|
async def _active_reservations(session: AsyncSession) -> int:
|
||||||
|
rows = await session.exec(
|
||||||
|
select(ReservationRelease).where(col(ReservationRelease.status) == "active")
|
||||||
|
)
|
||||||
|
return len(rows.all())
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_corrupt_revert_still_decrements_request_count(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
revert_pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, 3_000, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
await integration_session.exec( # type: ignore[call-overload]
|
||||||
|
update(ApiKey)
|
||||||
|
.where(col(ApiKey.hashed_key) == key_hash)
|
||||||
|
.values(reserved_balance=0)
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
assert await revert_pay_for_request(
|
||||||
|
key,
|
||||||
|
integration_session,
|
||||||
|
3_000,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
updated = await integration_session.get(ApiKey, key_hash)
|
||||||
|
record = await integration_session.get(ReservationRelease, reservation.release_id)
|
||||||
|
assert updated is not None
|
||||||
|
assert updated.total_requests == 0
|
||||||
|
assert updated.reserved_balance == 0
|
||||||
|
assert record is not None and record.status == "released"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_exact_cost_branch_charges_the_reservation_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == cost
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
assert key.total_requests == 1
|
||||||
|
assert await _active_reservations(integration_session) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_max_cost_branch_charges_the_reservation_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""No token pricing configured -> flat MaxCostData charge of the reservation."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 2_500
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"routstr.auth.calculate_cost",
|
||||||
|
return_value=_cost_data(cost, cls=MaxCostData),
|
||||||
|
):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == cost
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_underrun_branch_refunds_the_unused_reservation(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved = 5_000
|
||||||
|
actual = 1_200
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(actual)):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == actual
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.total_spent == actual, "user must pay the real cost, not the reservation"
|
||||||
|
assert key.balance == 10_000 - actual
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_overrun_branch_charges_full_cost_when_balance_is_free(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved = 1_000
|
||||||
|
actual = 1_400
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(actual)):
|
||||||
|
result = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result["charged_msats"] == actual
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.total_spent == actual
|
||||||
|
assert key.balance == 10_000 - actual
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_zero_cost_response_is_free_and_releases_the_reservation(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""An empty/unusable upstream response must cost the user nothing."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved = 4_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(0)):
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
{"model": "test-model"},
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
await integration_session.refresh(key)
|
||||||
|
assert key.balance == 10_000
|
||||||
|
assert key.total_spent == 0
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
assert await _active_reservations(integration_session) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_cost_error_releases_the_reservation_without_charging(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
reserved = 4_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, reserved, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"routstr.auth.calculate_cost",
|
||||||
|
return_value=CostDataError(message="no pricing", code="pricing_error"),
|
||||||
|
):
|
||||||
|
cost = await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
reserved,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
assert cost["charged_msats"] == 0
|
||||||
|
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.balance == 10_000, "a pricing failure must not charge the user"
|
||||||
|
assert key.total_spent == 0
|
||||||
|
assert key.reserved_balance == 0, "funds must not stay locked"
|
||||||
|
assert await _active_reservations(integration_session) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_repeated_finalization_charges_only_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A retried finalizer (proxy retry, duplicate stream end) must not double-bill."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
results = []
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
for _ in range(3):
|
||||||
|
# A declined re-charge rolls its session back, so re-load the key
|
||||||
|
# the way a fresh request would instead of reusing a stale instance.
|
||||||
|
integration_session.expunge_all()
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
results.append(
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Only the first finalization debits; duplicates report a zero charge.
|
||||||
|
assert [r["charged_msats"] for r in results] == [cost, 0, 0]
|
||||||
|
|
||||||
|
integration_session.expunge_all()
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.total_spent == cost, f"charged {key.total_spent} for one request"
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_concurrent_duplicate_finalization_charges_only_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
) -> None:
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
)
|
||||||
|
from routstr.core.db import create_session
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
async with create_session() as session:
|
||||||
|
key_hash = await _new_key(session, balance=10_000)
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
await pay_for_request(key, cost, session)
|
||||||
|
reservation = await get_reservation_snapshot(key, session)
|
||||||
|
|
||||||
|
async def finalize() -> None:
|
||||||
|
async with create_session() as session:
|
||||||
|
fresh = await session.get(ApiKey, key_hash)
|
||||||
|
assert fresh is not None
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
fresh,
|
||||||
|
_response(),
|
||||||
|
session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
await asyncio.gather(finalize(), finalize(), finalize())
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
key = await session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_release_after_charge_does_not_credit_the_user_back(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A late cleanup path must not turn a charged request into a free one."""
|
||||||
|
from routstr.auth import (
|
||||||
|
adjust_payment_for_tokens,
|
||||||
|
get_reservation_snapshot,
|
||||||
|
pay_for_request,
|
||||||
|
release_reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
cost = 3_000
|
||||||
|
key_hash = await _new_key(integration_session, balance=10_000)
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
|
||||||
|
await pay_for_request(key, cost, integration_session)
|
||||||
|
reservation = await get_reservation_snapshot(key, integration_session)
|
||||||
|
|
||||||
|
with patch("routstr.auth.calculate_cost", return_value=_cost_data(cost)):
|
||||||
|
await adjust_payment_for_tokens(
|
||||||
|
key,
|
||||||
|
_response(),
|
||||||
|
integration_session,
|
||||||
|
cost,
|
||||||
|
reservation_snapshot=reservation,
|
||||||
|
)
|
||||||
|
|
||||||
|
released = await release_reservation(reservation, integration_session, cost)
|
||||||
|
assert released is False, "a charged reservation must not be releasable"
|
||||||
|
|
||||||
|
key = await integration_session.get(ApiKey, key_hash)
|
||||||
|
assert key is not None
|
||||||
|
assert key.balance == 10_000 - cost
|
||||||
|
assert key.total_spent == cost
|
||||||
|
assert key.reserved_balance == 0
|
||||||
@@ -23,6 +23,7 @@ from routstr.upstream.auto_topup import (
|
|||||||
_ppq_request_id,
|
_ppq_request_id,
|
||||||
_ppq_spent_last_24h_usd,
|
_ppq_spent_last_24h_usd,
|
||||||
_ppq_state_id_for_provider,
|
_ppq_state_id_for_provider,
|
||||||
|
_reconcile_ppq_state,
|
||||||
_record_ppq_invoice,
|
_record_ppq_invoice,
|
||||||
_set_ppq_state_terminal,
|
_set_ppq_state_terminal,
|
||||||
get_ppq_auto_topup_state,
|
get_ppq_auto_topup_state,
|
||||||
@@ -35,6 +36,8 @@ pytestmark = pytest.mark.asyncio
|
|||||||
def _row(provider_id: int = 1) -> MagicMock:
|
def _row(provider_id: int = 1) -> MagicMock:
|
||||||
row = MagicMock()
|
row = MagicMock()
|
||||||
row.id = provider_id
|
row.id = provider_id
|
||||||
|
row.api_key = "secret"
|
||||||
|
row.provider_settings = None
|
||||||
return row
|
return row
|
||||||
|
|
||||||
|
|
||||||
@@ -111,12 +114,35 @@ async def test_claim_is_reusable_once_the_previous_attempt_finished(
|
|||||||
await _seed_provider()
|
await _seed_provider()
|
||||||
first = await _claim_ppq_topup(_row())
|
first = await _claim_ppq_topup(_row())
|
||||||
assert first is not None
|
assert first is not None
|
||||||
assert await _set_ppq_state_terminal(_row(), first, collected=True, swept=False)
|
assert await _set_ppq_state_terminal(_row(), first, collected=False, swept=True)
|
||||||
|
|
||||||
second = await _claim_ppq_topup(_row())
|
second = await _claim_ppq_topup(_row())
|
||||||
assert second is not None and second != first
|
assert second is not None and second != first
|
||||||
|
|
||||||
|
|
||||||
|
async def test_settled_claim_suppresses_immediate_duplicate(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
await _seed_provider()
|
||||||
|
row = _row()
|
||||||
|
operation_id = await _claim_ppq_topup(row)
|
||||||
|
assert operation_id is not None
|
||||||
|
assert await _set_ppq_state_terminal(row, operation_id, collected=True, swept=False)
|
||||||
|
|
||||||
|
assert await _reconcile_ppq_state(row, provider=None) is True
|
||||||
|
assert await _claim_ppq_topup(row) is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_claim_rejects_stale_provider_configuration(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
await _seed_provider()
|
||||||
|
stale = _row()
|
||||||
|
stale.provider_settings = '{"auto_topup":true}'
|
||||||
|
|
||||||
|
assert await _claim_ppq_topup(stale) is None
|
||||||
|
|
||||||
|
|
||||||
async def test_recording_the_invoice_moves_the_claim_in_flight(
|
async def test_recording_the_invoice_moves_the_claim_in_flight(
|
||||||
patched_db_engine: Any,
|
patched_db_engine: Any,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -287,7 +313,13 @@ async def test_ppq_payment_audit_row_is_visible_and_survives_next_claim(
|
|||||||
assert audit["collected"] is True
|
assert audit["collected"] is True
|
||||||
assert "lnbc-secret-invoice" not in audit["token"]
|
assert "lnbc-secret-invoice" not in audit["token"]
|
||||||
|
|
||||||
# Reusing the deterministic claim lock must not overwrite history.
|
assert await _claim_ppq_topup(_row()) is None
|
||||||
|
async with create_session() as session:
|
||||||
|
state = await session.get(CashuTransaction, _ppq_state_id_for_provider(1))
|
||||||
|
assert state is not None
|
||||||
|
state.created_at = int(time.time()) - 301
|
||||||
|
session.add(state)
|
||||||
|
await session.commit()
|
||||||
assert await _claim_ppq_topup(_row()) is not None
|
assert await _claim_ppq_topup(_row()) is not None
|
||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
assert await session.get(CashuTransaction, audit["id"]) is not None
|
assert await session.get(CashuTransaction, audit["id"]) is not None
|
||||||
|
|||||||
@@ -686,7 +686,7 @@ async def test_no_database_changes_during_provider_operations(
|
|||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_admin_routstr_topup_retries_transient_upstream_failure(
|
async def test_admin_routstr_topup_does_not_duplicate_invoice_on_upstream_failure(
|
||||||
integration_client: AsyncClient,
|
integration_client: AsyncClient,
|
||||||
integration_session: Any,
|
integration_session: Any,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -739,16 +739,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure(
|
|||||||
assert json["api_key"] == "sk-upstream-test"
|
assert json["api_key"] == "sk-upstream-test"
|
||||||
assert headers["Authorization"] == "Bearer sk-upstream-test"
|
assert headers["Authorization"] == "Bearer sk-upstream-test"
|
||||||
|
|
||||||
if self.calls == 1:
|
return MockResponse(500, {"detail": "ambiguous upstream failure"})
|
||||||
return MockResponse(500, {"detail": "warmup failure"})
|
|
||||||
|
|
||||||
return MockResponse(
|
|
||||||
200,
|
|
||||||
{
|
|
||||||
"bolt11": "lnbc1testinvoice",
|
|
||||||
"invoice_id": "invoice-123",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
mock_client = MockAsyncClient()
|
mock_client = MockAsyncClient()
|
||||||
|
|
||||||
@@ -759,11 +750,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure(
|
|||||||
json={"amount": 10},
|
json={"amount": 10},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 500
|
||||||
data = response.json()
|
assert mock_client.calls == 1
|
||||||
assert data["ok"] is True
|
|
||||||
assert data["topup_data"]["payment_request"] == "lnbc1testinvoice"
|
|
||||||
assert data["topup_data"]["invoice_id"] == "invoice-123"
|
|
||||||
assert mock_client.calls == 2
|
|
||||||
finally:
|
finally:
|
||||||
admin_sessions.pop(admin_token, None)
|
admin_sessions.pop(admin_token, None)
|
||||||
|
|||||||
@@ -289,8 +289,115 @@ async def test_proxy_post_unauthorized_access(integration_client: AsyncClient) -
|
|||||||
assert response.status_code in [400, 401]
|
assert response.status_code in [400, 401]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"bad_path",
|
||||||
|
[
|
||||||
|
"internal/admin", # unknown endpoint, no API prefix
|
||||||
|
"v1/../admin", # traversal onto a sibling path
|
||||||
|
"v1//models", # duplicate separator
|
||||||
|
"%2e%2e/secret", # encoded dot segment
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_authenticated_post_to_unknown_path_is_rejected(
|
||||||
|
authenticated_client: AsyncClient, bad_path: str
|
||||||
|
) -> None:
|
||||||
|
"""An authenticated POST to an unknown/traversal path must be rejected at
|
||||||
|
the edge (404) and never forwarded — the provider credential must not reach
|
||||||
|
an endpoint the caller merely spelled into the URL. If the guard let it
|
||||||
|
through, forwarding would raise and this would not be a clean 404."""
|
||||||
|
with patch(
|
||||||
|
"routstr.upstream.base.BaseUpstreamProvider.forward_request",
|
||||||
|
AsyncMock(side_effect=AssertionError("must not forward unknown path")),
|
||||||
|
):
|
||||||
|
response = await authenticated_client.post(
|
||||||
|
f"/{bad_path}",
|
||||||
|
json={"model": "gpt-3.5-turbo", "messages": []},
|
||||||
|
)
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"bad_path",
|
||||||
|
[
|
||||||
|
# Unambiguously spelled, under a prefix the proxy serves, but not an
|
||||||
|
# endpoint it offers. These are real upstream routes that manage keys,
|
||||||
|
# org membership, and billing on the same origin as inference.
|
||||||
|
"v1/organization/api_keys",
|
||||||
|
"v1/api_keys",
|
||||||
|
"v1/billing/usage",
|
||||||
|
"v1/admin/keys",
|
||||||
|
# An id segment is honoured one level deep, and only where an endpoint
|
||||||
|
# takes one at all.
|
||||||
|
"v1/chat/completions/abc",
|
||||||
|
"v1/models/gpt-4/secret",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_known_prefix_does_not_carry_an_unknown_endpoint(
|
||||||
|
authenticated_client: AsyncClient, bad_path: str
|
||||||
|
) -> None:
|
||||||
|
"""A familiar prefix must not be a passport for the rest of the origin.
|
||||||
|
|
||||||
|
Nothing here is traversal-shaped, so the spelling screen lets it by; only
|
||||||
|
the endpoint allowlist stops it. Forwarding raises if the guard misses,
|
||||||
|
so a clean 404 also proves the credential never left."""
|
||||||
|
with patch(
|
||||||
|
"routstr.upstream.base.BaseUpstreamProvider.forward_request",
|
||||||
|
AsyncMock(side_effect=AssertionError("must not forward unknown endpoint")),
|
||||||
|
):
|
||||||
|
response = await authenticated_client.post(
|
||||||
|
f"/{bad_path}",
|
||||||
|
json={"model": "gpt-3.5-turbo", "messages": []},
|
||||||
|
)
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_unknown_path_is_rejected(integration_client: AsyncClient) -> None:
|
||||||
|
"""A GET to an unknown (no API prefix) path is rejected at the edge."""
|
||||||
|
response = await integration_client.get("/internal/admin")
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"bad_path",
|
||||||
|
[
|
||||||
|
# Percent-encoded so the traversal survives the client to the server,
|
||||||
|
# which decodes it to "v1/../admin" before routing.
|
||||||
|
"v1/%2e%2e/admin",
|
||||||
|
# Well-spelled, so only the endpoint allowlist can stop it.
|
||||||
|
"v1/organization/api_keys",
|
||||||
|
"anything/encrypted",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_ehbp_request_is_gated_by_the_endpoint_allowlist(
|
||||||
|
integration_client: AsyncClient, bad_path: str
|
||||||
|
) -> None:
|
||||||
|
"""EHBP hides the body from the proxy, not the destination.
|
||||||
|
|
||||||
|
The encrypted contract covers the request body; it says nothing about which
|
||||||
|
endpoint the provider credential gets spent against, so an EHBP request is
|
||||||
|
screened and allowlisted exactly like any other."""
|
||||||
|
with patch(
|
||||||
|
"routstr.proxy.forward_ehbp_request",
|
||||||
|
AsyncMock(side_effect=AssertionError("must not forward unknown endpoint")),
|
||||||
|
):
|
||||||
|
response = await integration_client.post(
|
||||||
|
f"/{bad_path}",
|
||||||
|
content=b"encrypted",
|
||||||
|
headers={
|
||||||
|
"ehbp-encapsulated-key": "x",
|
||||||
|
"x-routstr-model": "gpt-4",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ def _dead_key(created_at: int | None) -> ApiKey:
|
|||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_prunes_old_refunded_zero_key(patched_db_engine: None) -> None:
|
async def test_prunes_old_refunded_zero_key(patched_db_engine: None) -> None:
|
||||||
"""A funded-then-refunded key (0/0/0, NULL parent, old) is pruned."""
|
"""A funded-then-refunded zero-balance key with an old timestamp is pruned."""
|
||||||
key = _dead_key(LONG_AGO)
|
key = _dead_key(LONG_AGO)
|
||||||
async with create_session() as session:
|
async with create_session() as session:
|
||||||
session.add(key)
|
session.add(key)
|
||||||
@@ -98,34 +98,6 @@ async def test_used_key_never_pruned(patched_db_engine: None) -> None:
|
|||||||
assert await _exists(k.hashed_key)
|
assert await _exists(k.hashed_key)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_parent_and_child_keys_are_not_pruned(
|
|
||||||
patched_db_engine: None,
|
|
||||||
) -> None:
|
|
||||||
"""Pruning must not orphan child keys or delete valid children."""
|
|
||||||
parent = _dead_key(LONG_AGO)
|
|
||||||
child = ApiKey(
|
|
||||||
hashed_key=f"child_{uuid.uuid4().hex}",
|
|
||||||
balance=0,
|
|
||||||
reserved_balance=0,
|
|
||||||
total_spent=0,
|
|
||||||
total_requests=0,
|
|
||||||
created_at=LONG_AGO,
|
|
||||||
parent_key_hash=parent.hashed_key,
|
|
||||||
)
|
|
||||||
async with create_session() as session:
|
|
||||||
session.add(parent)
|
|
||||||
session.add(child)
|
|
||||||
await session.commit()
|
|
||||||
|
|
||||||
async with create_session() as session:
|
|
||||||
pruned = await prune_dead_api_keys(session, OLD)
|
|
||||||
|
|
||||||
assert pruned == 0
|
|
||||||
assert await _exists(parent.hashed_key)
|
|
||||||
assert await _exists(child.hashed_key)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
("status", "expires_at"),
|
("status", "expires_at"),
|
||||||
@@ -321,3 +293,33 @@ async def test_periodic_prune_disabled_returns_immediately(
|
|||||||
await auth.periodic_dead_key_prune()
|
await auth.periodic_dead_key_prune()
|
||||||
|
|
||||||
sleep_mock.assert_not_called()
|
sleep_mock.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_refund_claim_protects_key(patched_db_engine: None) -> None:
|
||||||
|
"""A key with a refund row is the audit anchor for its payout and holds a
|
||||||
|
non-null FK, so the janitor must leave it alone."""
|
||||||
|
from routstr.core.db import Refund
|
||||||
|
|
||||||
|
key = _dead_key(LONG_AGO)
|
||||||
|
async with create_session() as session:
|
||||||
|
session.add(key)
|
||||||
|
await session.commit()
|
||||||
|
session.add(
|
||||||
|
Refund(
|
||||||
|
api_key_hashed_key=key.hashed_key,
|
||||||
|
method="lightning",
|
||||||
|
destination="user@ln.example.com",
|
||||||
|
amount_msats=1000,
|
||||||
|
unit="sat",
|
||||||
|
mint_url="https://mint.example.com",
|
||||||
|
status="paid",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
async with create_session() as session:
|
||||||
|
pruned = await prune_dead_api_keys(session, OLD)
|
||||||
|
|
||||||
|
assert pruned == 0
|
||||||
|
assert await _exists(key.hashed_key)
|
||||||
|
|||||||
@@ -0,0 +1,891 @@
|
|||||||
|
"""Refund claim lifecycle against a real SQLite database.
|
||||||
|
|
||||||
|
Covers the guarantees the ``refunds`` table exists to provide: one open claim
|
||||||
|
per key, a persisted melt quote before the melt is dispatched, and a
|
||||||
|
reconciler that never restores a balance whose payout may have settled.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
|
from typing import Any, Awaitable, Callable
|
||||||
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
import pytest
|
||||||
|
from fastapi import HTTPException
|
||||||
|
from sqlmodel import select
|
||||||
|
|
||||||
|
from routstr import refund
|
||||||
|
from routstr.balance import RefundRequest, refund_wallet_endpoint
|
||||||
|
from routstr.core.db import (
|
||||||
|
ApiKey,
|
||||||
|
AsyncSession,
|
||||||
|
CashuTransaction,
|
||||||
|
Refund,
|
||||||
|
store_cashu_transaction_with_retry,
|
||||||
|
total_user_liability,
|
||||||
|
)
|
||||||
|
from routstr.payment.lnurl import LNURLError, MeltOutcomeAmbiguousError
|
||||||
|
|
||||||
|
KEY_HASH = "refundclaimkey"
|
||||||
|
ADDRESS = "user@ln.example.com"
|
||||||
|
BALANCE_MSATS = 5_000_000
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_key(
|
||||||
|
session: AsyncSession, *, balance: int = BALANCE_MSATS, address: str | None = None
|
||||||
|
) -> ApiKey:
|
||||||
|
key = ApiKey(hashed_key=KEY_HASH)
|
||||||
|
key.balance = balance
|
||||||
|
key.reserved_balance = 0
|
||||||
|
key.refund_currency = "sat"
|
||||||
|
key.refund_address = address
|
||||||
|
key.total_spent = 0
|
||||||
|
key.total_requests = 0
|
||||||
|
session.add(key)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(key)
|
||||||
|
return key
|
||||||
|
|
||||||
|
|
||||||
|
async def _load_key(session: AsyncSession) -> ApiKey:
|
||||||
|
key = await session.get(ApiKey, KEY_HASH)
|
||||||
|
assert key is not None
|
||||||
|
await session.refresh(key)
|
||||||
|
return key
|
||||||
|
|
||||||
|
|
||||||
|
async def _load_refund(session: AsyncSession, refund_id: str) -> Refund:
|
||||||
|
row = await session.get(Refund, refund_id)
|
||||||
|
assert row is not None
|
||||||
|
await session.refresh(row)
|
||||||
|
return row
|
||||||
|
|
||||||
|
|
||||||
|
async def _age_claim(session: AsyncSession, refund_id: str, seconds: int) -> None:
|
||||||
|
row = await _load_refund(session, refund_id)
|
||||||
|
row.claimed_at = int(time.time()) - seconds
|
||||||
|
row.created_at = int(time.time()) - seconds
|
||||||
|
row.updated_at = int(time.time()) - seconds
|
||||||
|
session.add(row)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def short_timeout() -> Any:
|
||||||
|
with patch.object(refund.settings, "refund_claim_timeout_seconds", 300):
|
||||||
|
yield
|
||||||
|
|
||||||
|
|
||||||
|
# --- exclusivity -----------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_second_claim_on_open_key_is_rejected(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
first = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
key = await _load_key(integration_session)
|
||||||
|
assert key.balance == 0
|
||||||
|
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
detail = exc_info.value.detail
|
||||||
|
assert exc_info.value.status_code == 409
|
||||||
|
assert isinstance(detail, dict)
|
||||||
|
assert detail["error"]["code"] == "refund_in_progress"
|
||||||
|
|
||||||
|
rows = (await integration_session.exec(select(Refund))).all()
|
||||||
|
assert [row.id for row in rows] == [first.id]
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_claim_from_second_session_hits_the_index(
|
||||||
|
integration_engine: Any, integration_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
|
||||||
|
async with AsyncSession(integration_engine, expire_on_commit=False) as other:
|
||||||
|
other_key = await _load_key(other)
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.open_claim(
|
||||||
|
other, other_key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
assert exc_info.value.status_code == 409
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_claim_rejects_stale_balance_snapshot(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
integration_session.expunge(key) # a stale, detached snapshot
|
||||||
|
key.balance = BALANCE_MSATS + 1
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
assert exc_info.value.status_code == 409
|
||||||
|
assert (await integration_session.exec(select(Refund))).all() == []
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_retry_after_failed_claim_pays_once(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
first = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
assert await refund.release(integration_session, first)
|
||||||
|
key = await _load_key(integration_session)
|
||||||
|
assert key.balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
second = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
assert second.id != first.id
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_release_after_settle_is_a_noop(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
assert await refund.settle(integration_session, claim, quote_id="q1")
|
||||||
|
assert not await refund.release(integration_session, claim)
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert row.status == "paid"
|
||||||
|
assert row.claimed_at is None
|
||||||
|
|
||||||
|
|
||||||
|
# --- execute ---------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _lnurl_stub(
|
||||||
|
outcome: BaseException | None = None,
|
||||||
|
*,
|
||||||
|
quoted: bool = True,
|
||||||
|
) -> Callable[..., Awaitable[int]]:
|
||||||
|
async def send(
|
||||||
|
amount: int,
|
||||||
|
unit: str,
|
||||||
|
mint: str,
|
||||||
|
address: str,
|
||||||
|
*,
|
||||||
|
on_melt_quote: Callable[[str, str], Awaitable[None]] | None = None,
|
||||||
|
) -> int:
|
||||||
|
if quoted and on_melt_quote is not None:
|
||||||
|
await on_melt_quote("quote-123", mint)
|
||||||
|
if outcome is not None:
|
||||||
|
raise outcome
|
||||||
|
return amount
|
||||||
|
|
||||||
|
return send
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_execute_persists_quote_before_melt_and_settles(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
seen: list[str | None] = []
|
||||||
|
|
||||||
|
async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int:
|
||||||
|
await on_melt_quote("quote-123", claim.mint_url)
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
seen.append(row.quote_id)
|
||||||
|
return 5000
|
||||||
|
|
||||||
|
with patch("routstr.refund.send_to_lnurl", send):
|
||||||
|
body = await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
assert seen == ["quote-123"], "quote must be on disk before the melt runs"
|
||||||
|
assert body["status"] == "paid"
|
||||||
|
assert body["recipient"] == ADDRESS
|
||||||
|
assert body["sats"] == "5000"
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert (row.status, row.quote_id, row.claimed_at) == ("paid", "quote-123", None)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_execute_ambiguous_holds_claim_with_quote(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.send_to_lnurl", _lnurl_stub(MeltOutcomeAmbiguousError("?"))
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
assert exc_info.value.status_code == 502
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert (row.status, row.quote_id, row.claimed_at) == (
|
||||||
|
"ambiguous",
|
||||||
|
"quote-123",
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_execute_clean_failure_restores_balance(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.send_to_lnurl",
|
||||||
|
_lnurl_stub(LNURLError("limits"), quoted=False),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
assert exc_info.value.status_code == 500
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert row.status == "failed"
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_execute_failure_after_quote_withholds_balance(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""The mint may have paid the quote, so a later local failure must not restore."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
with patch("routstr.refund.send_to_lnurl", _lnurl_stub(RuntimeError("local"))):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
assert exc_info.value.status_code == 502
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert (row.status, row.quote_id) == ("ambiguous", "quote-123")
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_execute_records_mint_that_issued_the_quote(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""Mint fallback must leave reconciliation pointed at the issuing mint."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
fallback_mint = "https://fallback.mint.example"
|
||||||
|
|
||||||
|
async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int:
|
||||||
|
await on_melt_quote("quote-fallback", fallback_mint)
|
||||||
|
raise MeltOutcomeAmbiguousError("unknown")
|
||||||
|
|
||||||
|
with patch("routstr.refund.send_to_lnurl", send):
|
||||||
|
with pytest.raises(MeltOutcomeAmbiguousError):
|
||||||
|
await refund._pay_lightning(integration_session, claim)
|
||||||
|
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert (row.mint_url, row.quote_id, row.status) == (
|
||||||
|
fallback_mint,
|
||||||
|
"quote-fallback",
|
||||||
|
"ambiguous",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_execute_aborts_melt_when_claim_was_released(
|
||||||
|
integration_engine: Any, integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""A reconciler that released the claim first must stop the melt."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
melted = False
|
||||||
|
|
||||||
|
async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int:
|
||||||
|
nonlocal melted
|
||||||
|
async with AsyncSession(integration_engine, expire_on_commit=False) as other:
|
||||||
|
await refund.release(other, await _load_refund(other, claim.id))
|
||||||
|
await on_melt_quote("quote-123", claim.mint_url)
|
||||||
|
melted = True
|
||||||
|
return 5000
|
||||||
|
|
||||||
|
with patch("routstr.refund.send_to_lnurl", send):
|
||||||
|
with pytest.raises(HTTPException):
|
||||||
|
await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
assert melted is False
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
|
||||||
|
# --- reconciler ------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def _open_ambiguous(session: AsyncSession, quote_id: str | None) -> Refund:
|
||||||
|
key = await _seed_key(session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
await refund.hold(session, claim, quote_id)
|
||||||
|
return claim
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("mint_status", "expected_status", "expected_balance"),
|
||||||
|
[
|
||||||
|
("paid", "paid", 0),
|
||||||
|
("unpaid", "failed", BALANCE_MSATS),
|
||||||
|
("pending", "ambiguous", 0),
|
||||||
|
("unknown", "ambiguous", 0),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_reconcile_ambiguous_claims(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
short_timeout: None,
|
||||||
|
mint_status: str,
|
||||||
|
expected_status: str,
|
||||||
|
expected_balance: int,
|
||||||
|
) -> None:
|
||||||
|
claim = await _open_ambiguous(integration_session, "quote-123")
|
||||||
|
await _age_claim(integration_session, claim.id, 600)
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.check_bolt11_payment_status",
|
||||||
|
AsyncMock(return_value=mint_status),
|
||||||
|
) as check:
|
||||||
|
await refund.reconcile_once()
|
||||||
|
|
||||||
|
check.assert_awaited_once_with(claim.mint_url, "sat", "quote-123")
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert row.status == expected_status
|
||||||
|
assert (await _load_key(integration_session)).balance == expected_balance
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconcile_credits_balance_once_across_passes(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
|
||||||
|
) -> None:
|
||||||
|
claim = await _open_ambiguous(integration_session, "quote-123")
|
||||||
|
await _age_claim(integration_session, claim.id, 600)
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid")
|
||||||
|
):
|
||||||
|
await refund.reconcile_once()
|
||||||
|
await refund.reconcile_once()
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconcile_waits_before_trusting_recent_unpaid(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
|
||||||
|
) -> None:
|
||||||
|
"""A just-dispatched melt can report unpaid before it turns pending, so an
|
||||||
|
unpaid quote is only final once the claim has been quiet for a timeout."""
|
||||||
|
claim = await _open_ambiguous(integration_session, "quote-123")
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid")
|
||||||
|
):
|
||||||
|
await refund.reconcile_once()
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert row.status == "ambiguous"
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
await _age_claim(integration_session, claim.id, 600)
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid")
|
||||||
|
):
|
||||||
|
await refund.reconcile_once()
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert row.status == "failed"
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconcile_keeps_claim_that_gained_quote_mid_pass(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
|
||||||
|
) -> None:
|
||||||
|
"""The payout records its quote between the reconciler's read and its
|
||||||
|
release: the stale ``quote_id is None`` must not restore the balance."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
await _age_claim(integration_session, claim.id, 600)
|
||||||
|
|
||||||
|
real_lease = refund._lease
|
||||||
|
|
||||||
|
async def lease_then_quote(refund_id: str, now: int, cutoff: int) -> bool:
|
||||||
|
leased = await real_lease(refund_id, now, cutoff)
|
||||||
|
await refund.record_quote(claim, "late-quote", claim.mint_url)
|
||||||
|
return leased
|
||||||
|
|
||||||
|
with patch("routstr.refund._lease", lease_then_quote):
|
||||||
|
await refund.reconcile_once()
|
||||||
|
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert row.status == "pending"
|
||||||
|
assert row.quote_id == "late-quote"
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconcile_leaves_fresh_pending_claim_alone(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
with patch("routstr.refund.check_bolt11_payment_status", AsyncMock()) as check:
|
||||||
|
await refund.reconcile_once()
|
||||||
|
check.assert_not_awaited()
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert row.status == "pending"
|
||||||
|
assert row.claimed_at is not None
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconcile_releases_expired_claim_without_quote(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
|
||||||
|
) -> None:
|
||||||
|
"""No quote on disk means the mint was never asked to pay."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
await _age_claim(integration_session, claim.id, 600)
|
||||||
|
with patch("routstr.refund.check_bolt11_payment_status", AsyncMock()) as check:
|
||||||
|
await refund.reconcile_once()
|
||||||
|
check.assert_not_awaited()
|
||||||
|
assert (await _load_refund(integration_session, claim.id)).status == "failed"
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconcile_queries_mint_for_crashed_claim_with_quote(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
|
||||||
|
) -> None:
|
||||||
|
"""Crash after the quote was persisted: the mint decides, not the timeout."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
await refund.record_quote(claim, "quote-crash", claim.mint_url)
|
||||||
|
await _age_claim(integration_session, claim.id, 600)
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="paid")
|
||||||
|
) as check:
|
||||||
|
await refund.reconcile_once()
|
||||||
|
check.assert_awaited_once_with(claim.mint_url, "sat", "quote-crash")
|
||||||
|
assert (await _load_refund(integration_session, claim.id)).status == "paid"
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconcile_marks_expired_cashu_claim_stuck(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
await _age_claim(integration_session, claim.id, 600)
|
||||||
|
with patch("routstr.refund.logger") as log:
|
||||||
|
await refund.reconcile_once()
|
||||||
|
await refund.reconcile_once()
|
||||||
|
assert log.critical.call_count == 1
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert (row.status, row.claimed_at) == ("stuck", None)
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
# A stuck claim is closed, so the key is not permanently locked out.
|
||||||
|
key = await _load_key(integration_session)
|
||||||
|
key.balance = 1000
|
||||||
|
integration_session.add(key)
|
||||||
|
await integration_session.commit()
|
||||||
|
await refund.open_claim(integration_session, key, method="cashu", destination=None)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconcile_survives_one_failing_row(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None, short_timeout: None
|
||||||
|
) -> None:
|
||||||
|
claim = await _open_ambiguous(integration_session, "quote-123")
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.check_bolt11_payment_status",
|
||||||
|
AsyncMock(side_effect=RuntimeError("mint down")),
|
||||||
|
):
|
||||||
|
await refund.reconcile_once()
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert row.status == "ambiguous"
|
||||||
|
assert row.claimed_at is not None, "lease is kept until the next pass"
|
||||||
|
|
||||||
|
|
||||||
|
# --- endpoint --------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_endpoint_uses_requested_address_over_persisted(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
await _seed_key(integration_session, address="stored@ln.example.com")
|
||||||
|
send = AsyncMock(side_effect=_lnurl_stub())
|
||||||
|
with (
|
||||||
|
patch("routstr.refund.get_lnurl_data", AsyncMock()) as resolve,
|
||||||
|
patch("routstr.refund.send_to_lnurl", send),
|
||||||
|
):
|
||||||
|
body = await refund_wallet_endpoint(
|
||||||
|
refund_request=RefundRequest(lightning_address=ADDRESS),
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
resolve.assert_awaited_once_with(ADDRESS)
|
||||||
|
assert isinstance(body, dict)
|
||||||
|
assert body["recipient"] == ADDRESS
|
||||||
|
assert send.await_args is not None
|
||||||
|
assert send.await_args.args[3] == ADDRESS
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_endpoint_rejects_bad_address_without_debit(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
await _seed_key(integration_session)
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.get_lnurl_data", AsyncMock(side_effect=LNURLError("nope"))
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund_wallet_endpoint(
|
||||||
|
refund_request=RefundRequest(lightning_address="bad@example"),
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert exc_info.value.status_code == 400
|
||||||
|
assert (await integration_session.exec(select(Refund))).all() == []
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_endpoint_replays_paid_lightning_refund_on_empty_balance(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
await _seed_key(integration_session, address=ADDRESS)
|
||||||
|
with (
|
||||||
|
patch("routstr.refund.get_lnurl_data", AsyncMock()),
|
||||||
|
patch("routstr.refund.send_to_lnurl", _lnurl_stub()),
|
||||||
|
):
|
||||||
|
first = await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
second = await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert isinstance(first, dict) and isinstance(second, dict)
|
||||||
|
assert second["refund_id"] == first["refund_id"]
|
||||||
|
assert second["status"] == "paid"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_endpoint_refund_while_ambiguous_returns_409(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
await _open_ambiguous(integration_session, "quote-123")
|
||||||
|
key = await _load_key(integration_session)
|
||||||
|
key.balance = 2_000_000 # topped up while the melt is unresolved
|
||||||
|
integration_session.add(key)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
send = AsyncMock()
|
||||||
|
with (
|
||||||
|
patch("routstr.refund.get_lnurl_data", AsyncMock()),
|
||||||
|
patch("routstr.refund.send_to_lnurl", send),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund_wallet_endpoint(
|
||||||
|
refund_request=RefundRequest(lightning_address=ADDRESS),
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert exc_info.value.status_code == 409
|
||||||
|
send.assert_not_awaited()
|
||||||
|
assert (await _load_key(integration_session)).balance == 2_000_000
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_endpoint_zero_balance_with_open_claim_returns_409(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""A prior refund debited the balance and is still settling: the retry must
|
||||||
|
report refund_in_progress (409), not "no balance to refund" (400)."""
|
||||||
|
await _open_ambiguous(integration_session, "quote-123")
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
send = AsyncMock()
|
||||||
|
with (
|
||||||
|
patch("routstr.refund.get_lnurl_data", AsyncMock()),
|
||||||
|
patch("routstr.refund.send_to_lnurl", send),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert exc_info.value.status_code == 409
|
||||||
|
assert isinstance(exc_info.value.detail, dict)
|
||||||
|
assert exc_info.value.detail["error"]["code"] == "refund_in_progress"
|
||||||
|
send.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_cashu_token_survives_failed_ledger_write(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""The token is issued once the mint signs it; a failed cashu_transactions
|
||||||
|
insert must neither fail the request nor release the balance, and a retry
|
||||||
|
must replay the token from the claim row."""
|
||||||
|
await _seed_key(integration_session)
|
||||||
|
with (
|
||||||
|
patch("routstr.refund.send_token", AsyncMock(return_value="cashuAtoken")),
|
||||||
|
patch("routstr.refund.token_mint_url", lambda token, mint: mint),
|
||||||
|
patch(
|
||||||
|
"routstr.refund.store_cashu_transaction",
|
||||||
|
AsyncMock(side_effect=RuntimeError("db down")),
|
||||||
|
),
|
||||||
|
):
|
||||||
|
first = await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert isinstance(first, dict)
|
||||||
|
assert (first["token"], first["status"]) == ("cashuAtoken", "paid")
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
second = await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert isinstance(second, dict)
|
||||||
|
assert second["refund_id"] == first["refund_id"]
|
||||||
|
assert second["token"] == "cashuAtoken"
|
||||||
|
|
||||||
|
|
||||||
|
async def _topup(session: AsyncSession, amount: int = BALANCE_MSATS) -> None:
|
||||||
|
key = await _load_key(session)
|
||||||
|
key.balance = amount
|
||||||
|
session.add(key)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def _refund_cashu(
|
||||||
|
session: AsyncSession, token: str, *, ledger: bool
|
||||||
|
) -> dict[str, str]:
|
||||||
|
store = (
|
||||||
|
AsyncMock(side_effect=RuntimeError("db down"))
|
||||||
|
if not ledger
|
||||||
|
else store_cashu_transaction_with_retry
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch("routstr.refund.send_token", AsyncMock(return_value=token)),
|
||||||
|
patch("routstr.refund.token_mint_url", lambda t, mint: mint),
|
||||||
|
patch("routstr.refund.store_cashu_transaction", store),
|
||||||
|
):
|
||||||
|
body = await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=session,
|
||||||
|
)
|
||||||
|
assert isinstance(body, dict)
|
||||||
|
return body
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_replay_prefers_the_token_of_the_newest_paid_claim(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""An older ledger row must not answer for a newer claim whose write failed."""
|
||||||
|
await _seed_key(integration_session)
|
||||||
|
await _refund_cashu(integration_session, "cashuAold", ledger=True)
|
||||||
|
await _topup(integration_session)
|
||||||
|
newest = await _refund_cashu(integration_session, "cashuAnew", ledger=False)
|
||||||
|
|
||||||
|
replay = await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert isinstance(replay, dict)
|
||||||
|
assert replay["token"] == "cashuAnew"
|
||||||
|
assert replay["refund_id"] == newest["refund_id"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_swept_older_ledger_row_does_not_reject_the_newest_claim(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
await _seed_key(integration_session)
|
||||||
|
await _refund_cashu(integration_session, "cashuAold", ledger=True)
|
||||||
|
await _topup(integration_session)
|
||||||
|
await _refund_cashu(integration_session, "cashuAnew", ledger=False)
|
||||||
|
|
||||||
|
result = await integration_session.exec(
|
||||||
|
select(CashuTransaction).where(CashuTransaction.token == "cashuAold")
|
||||||
|
)
|
||||||
|
old_tx = result.one()
|
||||||
|
old_tx.swept = True
|
||||||
|
integration_session.add(old_tx)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
replay = await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert isinstance(replay, dict)
|
||||||
|
assert replay["token"] == "cashuAnew"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_open_claim_is_reported_over_an_older_paid_claim(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""A paid claim from a previous cycle must not be replayed as the outcome
|
||||||
|
of the claim that is still settling."""
|
||||||
|
await _seed_key(integration_session, address=ADDRESS)
|
||||||
|
with (
|
||||||
|
patch("routstr.refund.get_lnurl_data", AsyncMock()),
|
||||||
|
patch("routstr.refund.send_to_lnurl", _lnurl_stub()),
|
||||||
|
):
|
||||||
|
paid = await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert isinstance(paid, dict)
|
||||||
|
await _topup(integration_session)
|
||||||
|
key = await _load_key(integration_session)
|
||||||
|
open_claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
await refund.hold(integration_session, open_claim, "quote-open")
|
||||||
|
|
||||||
|
validate = AsyncMock()
|
||||||
|
with patch("routstr.refund.get_lnurl_data", validate):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund_wallet_endpoint(
|
||||||
|
refund_request=RefundRequest(lightning_address=ADDRESS),
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert exc_info.value.status_code == 409
|
||||||
|
assert isinstance(exc_info.value.detail, dict)
|
||||||
|
error = exc_info.value.detail["error"]
|
||||||
|
assert (error["refund_id"], error["status"]) == (open_claim.id, "ambiguous")
|
||||||
|
# A drained key must not be able to drive outbound destination lookups.
|
||||||
|
validate.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stuck_claim_is_reported_instead_of_no_balance(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
await refund._close(integration_session, claim, status="stuck")
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund_wallet_endpoint(
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert exc_info.value.status_code == 409
|
||||||
|
assert isinstance(exc_info.value.detail, dict)
|
||||||
|
error = exc_info.value.detail["error"]
|
||||||
|
assert (error["code"], error["refund_id"]) == ("refund_unresolved", claim.id)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("status", "counted"),
|
||||||
|
[
|
||||||
|
("pending", True),
|
||||||
|
("ambiguous", True),
|
||||||
|
("stuck", True),
|
||||||
|
("paid", False),
|
||||||
|
("failed", False),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_liability_covers_claims_until_they_resolve(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
status: str,
|
||||||
|
counted: bool,
|
||||||
|
) -> None:
|
||||||
|
"""Money in flight is still owed to the customer; an owner payout that read
|
||||||
|
only key balances could spend its backing."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
before = await total_user_liability(integration_session)
|
||||||
|
assert before == BALANCE_MSATS
|
||||||
|
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
assert await total_user_liability(integration_session) == BALANCE_MSATS
|
||||||
|
|
||||||
|
await refund._close(integration_session, claim, status=status)
|
||||||
|
await integration_session.commit()
|
||||||
|
expected = BALANCE_MSATS if counted else 0
|
||||||
|
assert await total_user_liability(integration_session) == expected
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_unreachable_destination_is_a_client_error() -> None:
|
||||||
|
with patch(
|
||||||
|
"routstr.refund.get_lnurl_data",
|
||||||
|
AsyncMock(side_effect=httpx.ConnectError("All connection attempts failed")),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.validate_lightning_destination(ADDRESS)
|
||||||
|
assert exc_info.value.status_code == 400
|
||||||
@@ -0,0 +1,368 @@
|
|||||||
|
"""Refund payouts that finish after the claim row stopped cooperating.
|
||||||
|
|
||||||
|
Each test pins one guarantee of the claim table that the happy path cannot
|
||||||
|
exercise: the payout side has authoritative knowledge of what the mint did,
|
||||||
|
and the claim row must end up agreeing with it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
|
from typing import Any
|
||||||
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi import HTTPException
|
||||||
|
from sqlalchemy import event
|
||||||
|
from sqlalchemy.exc import OperationalError
|
||||||
|
|
||||||
|
from routstr import refund
|
||||||
|
from routstr.balance import RefundRequest, refund_wallet_endpoint
|
||||||
|
from routstr.core.db import ApiKey, AsyncSession, Refund, total_user_liability
|
||||||
|
from routstr.payment.lnurl import MeltUnpaidError
|
||||||
|
|
||||||
|
KEY_HASH = "refundrecoverykey"
|
||||||
|
ADDRESS = "user@ln.example.com"
|
||||||
|
BALANCE_MSATS = 5_000_000
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_key(session: AsyncSession, *, address: str | None = None) -> ApiKey:
|
||||||
|
key = ApiKey(hashed_key=KEY_HASH)
|
||||||
|
key.balance = BALANCE_MSATS
|
||||||
|
key.reserved_balance = 0
|
||||||
|
key.refund_currency = "sat"
|
||||||
|
key.refund_address = address
|
||||||
|
key.total_spent = 0
|
||||||
|
key.total_requests = 0
|
||||||
|
session.add(key)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(key)
|
||||||
|
return key
|
||||||
|
|
||||||
|
|
||||||
|
async def _load_key(session: AsyncSession) -> ApiKey:
|
||||||
|
key = await session.get(ApiKey, KEY_HASH)
|
||||||
|
assert key is not None
|
||||||
|
await session.refresh(key)
|
||||||
|
return key
|
||||||
|
|
||||||
|
|
||||||
|
async def _load_refund(session: AsyncSession, refund_id: str) -> Refund:
|
||||||
|
row = await session.get(Refund, refund_id)
|
||||||
|
assert row is not None
|
||||||
|
await session.refresh(row)
|
||||||
|
return row
|
||||||
|
|
||||||
|
|
||||||
|
def _cashu_payout(token: str, send: Any | None = None) -> Any:
|
||||||
|
return (
|
||||||
|
patch("routstr.refund.send_token", send or AsyncMock(return_value=token)),
|
||||||
|
patch("routstr.refund.token_mint_url", lambda t, mint: mint),
|
||||||
|
patch("routstr.refund.store_cashu_transaction", AsyncMock()),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# --- cashu token issued after the reconciler gave up -----------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_cashu_token_issued_after_reconciler_marked_claim_stuck_settles_it(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""A token in the customer's hands must leave its claim ``paid``: a row
|
||||||
|
left ``stuck`` keeps the amount in liability forever and tells the operator
|
||||||
|
to reconcile a payout that already happened."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
|
||||||
|
async def slow_send_token(amount: int, unit: str, mint_url: str) -> str:
|
||||||
|
# The lease lapses while the mint is still working.
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
row.claimed_at = (row.claimed_at or 0) - 10_000
|
||||||
|
integration_session.add(row)
|
||||||
|
await integration_session.commit()
|
||||||
|
await refund.reconcile_once()
|
||||||
|
assert (await _load_refund(integration_session, claim.id)).status == "stuck"
|
||||||
|
return "cashuAlate"
|
||||||
|
|
||||||
|
send, mint, store = _cashu_payout("cashuAlate", slow_send_token)
|
||||||
|
with send, mint, store:
|
||||||
|
body = await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
assert (body["token"], body["status"]) == ("cashuAlate", "paid")
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert (row.status, row.token, row.claimed_at) == ("paid", "cashuAlate", None)
|
||||||
|
assert await total_user_liability(integration_session) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_cashu_payout_renews_lease_before_asking_the_mint(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""The reconciler leaves a claim alone while its lease is fresh, so the
|
||||||
|
payout renews it right before the slow step."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
row.claimed_at = (row.claimed_at or 0) - 10_000
|
||||||
|
integration_session.add(row)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
async def send_token(amount: int, unit: str, mint_url: str) -> str:
|
||||||
|
await refund.reconcile_once()
|
||||||
|
return "cashuAfresh"
|
||||||
|
|
||||||
|
send, mint, store = _cashu_payout("cashuAfresh", send_token)
|
||||||
|
with send, mint, store, patch("routstr.refund.logger") as log:
|
||||||
|
await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
log.critical.assert_not_called()
|
||||||
|
assert (await _load_refund(integration_session, claim.id)).status == "paid"
|
||||||
|
|
||||||
|
|
||||||
|
# --- cashu token issued, claim write failed --------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize("failure_point", ["execute", "autoflush", "commit"])
|
||||||
|
async def test_cashu_claim_write_failure_after_token_creation_withholds_balance(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
integration_engine: Any,
|
||||||
|
failure_point: str,
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
claim_id = claim.id
|
||||||
|
fired = False
|
||||||
|
armed = False
|
||||||
|
|
||||||
|
def fail_once(*args: Any) -> None:
|
||||||
|
nonlocal fired
|
||||||
|
statement = args[2] if failure_point != "commit" else "COMMIT"
|
||||||
|
if (
|
||||||
|
armed
|
||||||
|
and not fired
|
||||||
|
and (statement.startswith("UPDATE refunds") or statement == "COMMIT")
|
||||||
|
):
|
||||||
|
fired = True
|
||||||
|
raise OperationalError(statement, {}, Exception("database is locked"))
|
||||||
|
|
||||||
|
async def send_token(amount: int, unit: str, mint_url: str) -> str:
|
||||||
|
nonlocal armed
|
||||||
|
if failure_point == "autoflush":
|
||||||
|
# A pending ORM write makes SQLAlchemy invalidate the transaction
|
||||||
|
# and expire attached objects when the actual SQL execution fails.
|
||||||
|
claim.updated_at -= 1
|
||||||
|
armed = True
|
||||||
|
return "cashuAstranded"
|
||||||
|
|
||||||
|
event_name = "commit" if failure_point == "commit" else "before_cursor_execute"
|
||||||
|
event.listen(integration_engine.sync_engine, event_name, fail_once)
|
||||||
|
send, _, store = _cashu_payout("cashuAstranded", send_token)
|
||||||
|
try:
|
||||||
|
with (
|
||||||
|
send,
|
||||||
|
store,
|
||||||
|
patch(
|
||||||
|
"routstr.refund.token_mint_url", return_value="https://fallback.mint"
|
||||||
|
),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.execute(integration_session, claim)
|
||||||
|
finally:
|
||||||
|
event.remove(integration_engine.sync_engine, event_name, fail_once)
|
||||||
|
|
||||||
|
assert fired
|
||||||
|
assert exc_info.value.status_code == 502
|
||||||
|
row = await _load_refund(integration_session, claim_id)
|
||||||
|
assert (row.status, row.token, row.mint_url) == (
|
||||||
|
"ambiguous",
|
||||||
|
"cashuAstranded",
|
||||||
|
"https://fallback.mint",
|
||||||
|
)
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
assert await total_user_liability(integration_session) == BALANCE_MSATS
|
||||||
|
|
||||||
|
await refund.reconcile_once()
|
||||||
|
row = await _load_refund(integration_session, claim_id)
|
||||||
|
assert (row.status, row.token) == ("paid", "cashuAstranded")
|
||||||
|
assert await total_user_liability(integration_session) == 0
|
||||||
|
with patch("routstr.refund.send_token", AsyncMock()) as send_again:
|
||||||
|
replay = await refund_wallet_endpoint(
|
||||||
|
refund_request=RefundRequest(),
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert isinstance(replay, dict)
|
||||||
|
assert replay["token"] == "cashuAstranded"
|
||||||
|
send_again.assert_not_awaited()
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reconciler_settles_held_cashu_claim_that_carries_a_token(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""A held cashu claim whose row carries the token is a completed payout;
|
||||||
|
the reconciler closes it as paid instead of escalating it to stuck."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
await refund.hold(integration_session, claim, None, token="cashuAheld")
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
row.claimed_at = None
|
||||||
|
row.updated_at -= 10_000
|
||||||
|
integration_session.add(row)
|
||||||
|
await integration_session.commit()
|
||||||
|
|
||||||
|
with patch("routstr.refund.logger") as log:
|
||||||
|
await refund.reconcile_once()
|
||||||
|
|
||||||
|
log.critical.assert_not_called()
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert (row.status, row.token) == ("paid", "cashuAheld")
|
||||||
|
assert await total_user_liability(integration_session) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize("hold_before_lease", [True, False])
|
||||||
|
async def test_stale_cashu_reconciliation_preserves_newly_held_token(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
patched_db_engine: None,
|
||||||
|
integration_engine: Any,
|
||||||
|
hold_before_lease: bool,
|
||||||
|
) -> None:
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
claim_id = claim.id
|
||||||
|
now = int(time.time())
|
||||||
|
cutoff = now - refund.settings.refund_claim_timeout_seconds
|
||||||
|
claim.claimed_at = cutoff - 1
|
||||||
|
await integration_session.commit()
|
||||||
|
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||||
|
stale = await session.get(Refund, claim_id)
|
||||||
|
assert stale is not None and stale.token is None
|
||||||
|
|
||||||
|
if hold_before_lease:
|
||||||
|
await refund.hold(integration_session, claim, None, token="cashuAheld")
|
||||||
|
assert await refund._lease(claim_id, now, cutoff)
|
||||||
|
real_close = refund._close
|
||||||
|
|
||||||
|
async def hold_before_close(session: AsyncSession, row: Refund, **kw: Any) -> bool:
|
||||||
|
if not hold_before_lease and kw.get("status") == "stuck":
|
||||||
|
# Token arrives at the last moment, even after a potential reload.
|
||||||
|
await real_close(
|
||||||
|
integration_session, claim, status="ambiguous", token="cashuAheld"
|
||||||
|
)
|
||||||
|
await integration_session.commit()
|
||||||
|
return await real_close(session, row, **kw)
|
||||||
|
|
||||||
|
with patch.object(refund, "_close", hold_before_close):
|
||||||
|
await refund._reconcile(stale, now)
|
||||||
|
|
||||||
|
row = await _load_refund(integration_session, claim_id)
|
||||||
|
assert (row.status, row.token) == ("paid", "cashuAheld")
|
||||||
|
assert (await _load_key(integration_session)).balance == 0
|
||||||
|
assert await total_user_liability(integration_session) == 0
|
||||||
|
with patch("routstr.refund.send_token", AsyncMock()) as send_again:
|
||||||
|
replay = await refund_wallet_endpoint(
|
||||||
|
refund_request=RefundRequest(),
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
assert isinstance(replay, dict)
|
||||||
|
assert replay["token"] == "cashuAheld"
|
||||||
|
send_again.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
# --- mint proved the melt unpaid -------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mint_confirmed_unpaid_melt_restores_balance_immediately(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""The mint answering ``unpaid`` to the melt itself is proof no payment
|
||||||
|
happened, so the customer gets the balance back now, not after the
|
||||||
|
reconciler timeout."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="lightning", destination=ADDRESS
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int:
|
||||||
|
await on_melt_quote("quote-unpaid", claim.mint_url)
|
||||||
|
raise MeltUnpaidError("Cashu mint confirmed that the melt was unpaid")
|
||||||
|
|
||||||
|
with patch("routstr.refund.send_to_lnurl", send):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
assert exc_info.value.status_code == 503
|
||||||
|
row = await _load_refund(integration_session, claim.id)
|
||||||
|
assert (row.status, row.quote_id) == ("failed", "quote-unpaid")
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
|
|
||||||
|
|
||||||
|
# --- response reflects the persisted claim ---------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_response_status_reflects_persisted_claim(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""When ``settle`` closes nothing the response must not invent ``paid``."""
|
||||||
|
key = await _seed_key(integration_session)
|
||||||
|
claim = await refund.open_claim(
|
||||||
|
integration_session, key, method="cashu", destination=None
|
||||||
|
)
|
||||||
|
|
||||||
|
async def send_token(amount: int, unit: str, mint_url: str) -> str:
|
||||||
|
# Somebody closed the row as failed while the mint was working.
|
||||||
|
await refund._close(integration_session, claim, status="failed")
|
||||||
|
await integration_session.commit()
|
||||||
|
return "cashuAorphan"
|
||||||
|
|
||||||
|
send, mint, store = _cashu_payout("cashuAorphan", send_token)
|
||||||
|
with send, mint, store:
|
||||||
|
body = await refund.execute(integration_session, claim)
|
||||||
|
|
||||||
|
assert body["status"] == "failed"
|
||||||
|
assert (await _load_refund(integration_session, claim.id)).status == "failed"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_stored_refund_address_is_validated_before_debit(
|
||||||
|
integration_session: AsyncSession, patched_db_engine: None
|
||||||
|
) -> None:
|
||||||
|
"""A bad address stored on the key is a client error, not a payout failure."""
|
||||||
|
await _seed_key(integration_session, address="nobody@invalid.example")
|
||||||
|
send = AsyncMock()
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"routstr.refund.get_lnurl_data",
|
||||||
|
AsyncMock(side_effect=refund.LNURLError("no such user")),
|
||||||
|
),
|
||||||
|
patch("routstr.refund.send_to_lnurl", send),
|
||||||
|
):
|
||||||
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
|
await refund_wallet_endpoint(
|
||||||
|
refund_request=RefundRequest(),
|
||||||
|
authorization=f"Bearer sk-{KEY_HASH}",
|
||||||
|
x_cashu=None,
|
||||||
|
session=integration_session,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert exc_info.value.status_code == 400
|
||||||
|
send.assert_not_awaited()
|
||||||
|
assert (await _load_key(integration_session)).balance == BALANCE_MSATS
|
||||||
@@ -133,14 +133,11 @@ async def test_reserved_balance_with_successful_requests(
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_revert_with_zero_reserved_balance_is_noop(
|
async def test_revert_with_zero_reserved_balance_repairs_terminally(
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Test that revert_pay_for_request is a no-op when reserved_balance is 0.
|
"""Reverting after the aggregate was already zeroed must not drive it
|
||||||
|
negative: the corrupt durable reservation is released without subtraction."""
|
||||||
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 pay_for_request, 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]}"
|
unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}"
|
||||||
@@ -153,21 +150,24 @@ async def test_revert_with_zero_reserved_balance_is_noop(
|
|||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
await pay_for_request(test_key, 100, integration_session)
|
await pay_for_request(test_key, 100, integration_session)
|
||||||
test_key.reserved_balance = 0
|
test_key.reserved_balance = 0
|
||||||
|
test_key.total_requests = 0
|
||||||
integration_session.add(test_key)
|
integration_session.add(test_key)
|
||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
|
|
||||||
# A stale cleanup already released the aggregate reservation.
|
# A stale cleanup already released the aggregate reservation. The revert
|
||||||
|
# terminalizes the durable row (repair) without driving the aggregate
|
||||||
|
# negative.
|
||||||
result = await revert_pay_for_request(test_key, integration_session, 100)
|
result = await revert_pay_for_request(test_key, integration_session, 100)
|
||||||
|
|
||||||
await integration_session.refresh(test_key)
|
integration_session.expunge_all()
|
||||||
|
updated = await integration_session.get(ApiKey, unique_key)
|
||||||
|
assert updated is not None
|
||||||
|
|
||||||
assert result is False, "Revert should return False when reservation already released"
|
assert result is True, "Revert must terminalize the corrupt reservation"
|
||||||
assert test_key.reserved_balance == 0, (
|
assert updated.reserved_balance == 0, (
|
||||||
f"Reserved balance should remain 0, got: {test_key.reserved_balance}"
|
f"Reserved balance should remain 0, got: {updated.reserved_balance}"
|
||||||
)
|
|
||||||
assert test_key.total_requests == 1, (
|
|
||||||
f"Total requests should remain 1, got: {test_key.total_requests}"
|
|
||||||
)
|
)
|
||||||
|
assert updated.total_requests == 0
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -203,10 +203,11 @@ async def test_revert_with_sufficient_reserved_balance_succeeds(
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_revert_partial_reserved_balance_is_noop(
|
async def test_revert_partial_reserved_balance_repairs_terminally(
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Test that reverting more than the current reserved_balance is a no-op."""
|
"""Reverting more than the aggregate holds must not clamp or go negative:
|
||||||
|
the corrupt durable reservation is released without subtraction."""
|
||||||
from routstr.auth import pay_for_request, revert_pay_for_request
|
from routstr.auth import pay_for_request, revert_pay_for_request
|
||||||
|
|
||||||
unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}"
|
unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}"
|
||||||
@@ -223,18 +224,19 @@ async def test_revert_partial_reserved_balance_is_noop(
|
|||||||
integration_session.add(test_key)
|
integration_session.add(test_key)
|
||||||
await integration_session.commit()
|
await integration_session.commit()
|
||||||
|
|
||||||
# Try to revert 500 when only 50 is reserved — should be no-op
|
# Reverting 500 when only 50 is reserved cannot subtract; the corrupt
|
||||||
|
# reservation is terminalized and the aggregate left untouched.
|
||||||
result = await revert_pay_for_request(test_key, integration_session, 500)
|
result = await revert_pay_for_request(test_key, integration_session, 500)
|
||||||
|
|
||||||
await integration_session.refresh(test_key)
|
integration_session.expunge_all()
|
||||||
|
updated = await integration_session.get(ApiKey, unique_key)
|
||||||
|
assert updated is not None
|
||||||
|
|
||||||
assert result is False, "Revert should fail when cost > reserved_balance"
|
assert result is True, "Revert must terminalize the corrupt reservation"
|
||||||
assert test_key.reserved_balance == 50, (
|
assert updated.reserved_balance == 50, (
|
||||||
f"Reserved balance should stay at 50, got: {test_key.reserved_balance}"
|
f"Reserved balance should stay at 50, got: {updated.reserved_balance}"
|
||||||
)
|
|
||||||
assert test_key.total_requests == 1, (
|
|
||||||
f"Total requests should stay at 1, got: {test_key.total_requests}"
|
|
||||||
)
|
)
|
||||||
|
assert updated.total_requests == 0
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -265,9 +267,7 @@ async def test_double_revert_prevented(
|
|||||||
snapshot = await get_reservation_snapshot(test_key, integration_session)
|
snapshot = await get_reservation_snapshot(test_key, integration_session)
|
||||||
|
|
||||||
# First revert — should succeed
|
# First revert — should succeed
|
||||||
result1 = await revert_pay_for_request(
|
result1 = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
|
||||||
test_key, integration_session, 500, snapshot
|
|
||||||
)
|
|
||||||
await integration_session.refresh(test_key)
|
await integration_session.refresh(test_key)
|
||||||
|
|
||||||
assert result1 is True
|
assert result1 is True
|
||||||
@@ -275,9 +275,7 @@ async def test_double_revert_prevented(
|
|||||||
assert test_key.total_requests == 4
|
assert test_key.total_requests == 4
|
||||||
|
|
||||||
# Second revert of the same amount — should be no-op
|
# Second revert of the same amount — should be no-op
|
||||||
result2 = await revert_pay_for_request(
|
result2 = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
|
||||||
test_key, integration_session, 500, snapshot
|
|
||||||
)
|
|
||||||
await integration_session.refresh(test_key)
|
await integration_session.refresh(test_key)
|
||||||
|
|
||||||
assert result2 is False, "Second revert should be a no-op"
|
assert result2 is False, "Second revert should be a no-op"
|
||||||
@@ -319,9 +317,7 @@ async def test_sequential_reverts_never_go_negative(
|
|||||||
# Run 5 sequential reverts for the same 500 reservation
|
# Run 5 sequential reverts for the same 500 reservation
|
||||||
results = []
|
results = []
|
||||||
for _ in range(5):
|
for _ in range(5):
|
||||||
r = await revert_pay_for_request(
|
r = await revert_pay_for_request(test_key, integration_session, 500, snapshot)
|
||||||
test_key, integration_session, 500, snapshot
|
|
||||||
)
|
|
||||||
results.append(r)
|
results.append(r)
|
||||||
|
|
||||||
await integration_session.refresh(test_key)
|
await integration_session.refresh(test_key)
|
||||||
@@ -337,63 +333,3 @@ async def test_sequential_reverts_never_go_negative(
|
|||||||
assert test_key.reserved_balance >= 0, (
|
assert test_key.reserved_balance >= 0, (
|
||||||
f"Reserved balance went negative: {test_key.reserved_balance}"
|
f"Reserved balance went negative: {test_key.reserved_balance}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_child_key_revert_floor_guard(
|
|
||||||
integration_session: AsyncSession,
|
|
||||||
) -> None:
|
|
||||||
"""Test that child key reserved_balance also has floor guard on revert."""
|
|
||||||
from routstr.auth import (
|
|
||||||
get_reservation_snapshot,
|
|
||||||
pay_for_request,
|
|
||||||
revert_pay_for_request,
|
|
||||||
)
|
|
||||||
|
|
||||||
parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}"
|
|
||||||
child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}"
|
|
||||||
|
|
||||||
parent_key = ApiKey(
|
|
||||||
hashed_key=parent_key_hash,
|
|
||||||
balance=10000,
|
|
||||||
reserved_balance=0,
|
|
||||||
total_requests=2,
|
|
||||||
)
|
|
||||||
child_key = ApiKey(
|
|
||||||
hashed_key=child_key_hash,
|
|
||||||
balance=0,
|
|
||||||
reserved_balance=0,
|
|
||||||
total_requests=2,
|
|
||||||
parent_key_hash=parent_key_hash,
|
|
||||||
)
|
|
||||||
integration_session.add(parent_key)
|
|
||||||
integration_session.add(child_key)
|
|
||||||
await integration_session.commit()
|
|
||||||
await pay_for_request(child_key, 500, integration_session)
|
|
||||||
snapshot = await get_reservation_snapshot(child_key, integration_session)
|
|
||||||
|
|
||||||
# First revert succeeds
|
|
||||||
result1 = await revert_pay_for_request(
|
|
||||||
child_key, integration_session, 500, snapshot
|
|
||||||
)
|
|
||||||
await integration_session.refresh(parent_key)
|
|
||||||
await integration_session.refresh(child_key)
|
|
||||||
|
|
||||||
assert result1 is True
|
|
||||||
assert parent_key.reserved_balance == 0
|
|
||||||
assert child_key.reserved_balance == 0
|
|
||||||
|
|
||||||
# Second revert is a no-op for both parent and child
|
|
||||||
result2 = await revert_pay_for_request(
|
|
||||||
child_key, integration_session, 500, snapshot
|
|
||||||
)
|
|
||||||
await integration_session.refresh(parent_key)
|
|
||||||
await integration_session.refresh(child_key)
|
|
||||||
|
|
||||||
assert result2 is False
|
|
||||||
assert parent_key.reserved_balance == 0, (
|
|
||||||
f"Parent reserved_balance should stay 0, got: {parent_key.reserved_balance}"
|
|
||||||
)
|
|
||||||
assert child_key.reserved_balance == 0, (
|
|
||||||
f"Child reserved_balance should stay 0, got: {child_key.reserved_balance}"
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -0,0 +1,380 @@
|
|||||||
|
"""Real-database tests for the Routstr-to-Routstr auto top-up spend bound.
|
||||||
|
|
||||||
|
The bound has to survive a process restart and concurrent workers, so these
|
||||||
|
run against actual SQL instead of mocked sessions: an in-memory counter would
|
||||||
|
pass a mocked test and still let a non-crediting peer drain the wallet.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import importlib
|
||||||
|
import time
|
||||||
|
from contextlib import ExitStack
|
||||||
|
from typing import Any
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sqlmodel import select
|
||||||
|
|
||||||
|
from routstr.core.db import CashuTransaction, UpstreamProviderRow, create_session
|
||||||
|
from routstr.upstream import auto_topup as auto_topup_module
|
||||||
|
from routstr.upstream.auto_topup import (
|
||||||
|
ROUTSTR_MAX_DAILY_TOPUP_SATS,
|
||||||
|
ROUTSTR_MAX_TOPUP_FAILURES,
|
||||||
|
ROUTSTR_PHASE_BACKOFF,
|
||||||
|
ROUTSTR_PHASE_HALTED,
|
||||||
|
ROUTSTR_PHASE_SENT,
|
||||||
|
_check_and_topup,
|
||||||
|
_claim_routstr_topup,
|
||||||
|
_parse_routstr_request_id,
|
||||||
|
_persist_routstr_token_and_mark_sent,
|
||||||
|
_routstr_spent_last_24h_sats,
|
||||||
|
_routstr_state_id_for_provider,
|
||||||
|
get_routstr_auto_topup_state,
|
||||||
|
release_routstr_auto_topup_state,
|
||||||
|
)
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.asyncio
|
||||||
|
|
||||||
|
TOPUP_SATS = 50
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed_provider(provider_id: int = 1) -> UpstreamProviderRow:
|
||||||
|
import json
|
||||||
|
|
||||||
|
row = UpstreamProviderRow(
|
||||||
|
id=provider_id,
|
||||||
|
slug=f"peer-{provider_id}",
|
||||||
|
provider_type="routstr",
|
||||||
|
base_url="https://peer.test",
|
||||||
|
api_key="secret",
|
||||||
|
enabled=True,
|
||||||
|
provider_settings=json.dumps(
|
||||||
|
{
|
||||||
|
"auto_topup": True,
|
||||||
|
"topup_threshold": 1,
|
||||||
|
"topup_amount_limit": TOPUP_SATS,
|
||||||
|
"topup_mint_url": "https://mint.test",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
)
|
||||||
|
async with create_session() as session:
|
||||||
|
session.add(row)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(row)
|
||||||
|
return row
|
||||||
|
|
||||||
|
|
||||||
|
def _peer(balance: float, *, topup: object = None) -> MagicMock:
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_balance = AsyncMock(return_value=balance)
|
||||||
|
provider.topup = AsyncMock(return_value=topup or {"balance": balance})
|
||||||
|
return provider
|
||||||
|
|
||||||
|
|
||||||
|
def _patch_wallet(module: Any, peer: MagicMock, token: str) -> ExitStack:
|
||||||
|
stack = ExitStack()
|
||||||
|
stack.enter_context(
|
||||||
|
patch.object(module.RoutstrUpstreamProvider, "from_db_row", return_value=peer)
|
||||||
|
)
|
||||||
|
stack.enter_context(
|
||||||
|
patch.object(
|
||||||
|
module, "send_token_from_owner_locked", AsyncMock(return_value=token)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
stack.enter_context(
|
||||||
|
patch.object(module, "token_mint_url", return_value="https://mint.test")
|
||||||
|
)
|
||||||
|
return stack
|
||||||
|
|
||||||
|
|
||||||
|
async def _claim_state(provider_id: int = 1) -> CashuTransaction | None:
|
||||||
|
async with create_session() as session:
|
||||||
|
return await session.get(
|
||||||
|
CashuTransaction, _routstr_state_id_for_provider(provider_id)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _sent_tokens() -> list[CashuTransaction]:
|
||||||
|
async with create_session() as session:
|
||||||
|
return list(
|
||||||
|
(
|
||||||
|
await session.exec(
|
||||||
|
select(CashuTransaction).where(
|
||||||
|
CashuTransaction.source == "auto_topup"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_second_worker_cannot_claim_while_the_first_holds_one(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
assert await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) is not None
|
||||||
|
assert await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_token_and_sent_claim_roll_back_together_on_commit_failure(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
operation_id = await _claim_routstr_topup(row, expected_sats=TOPUP_SATS)
|
||||||
|
assert operation_id is not None
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sqlmodel.ext.asyncio.session.AsyncSession.commit",
|
||||||
|
new=AsyncMock(side_effect=RuntimeError("commit failed")),
|
||||||
|
):
|
||||||
|
with pytest.raises(RuntimeError, match="commit failed"):
|
||||||
|
await _persist_routstr_token_and_mark_sent(
|
||||||
|
row,
|
||||||
|
operation_id,
|
||||||
|
expected_sats=TOPUP_SATS,
|
||||||
|
token="cashu-token-atomic",
|
||||||
|
amount=TOPUP_SATS,
|
||||||
|
mint_url="https://mint.test",
|
||||||
|
)
|
||||||
|
|
||||||
|
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
|
||||||
|
assert claim is not None and claim.phase != ROUTSTR_PHASE_SENT
|
||||||
|
assert await _sent_tokens() == []
|
||||||
|
|
||||||
|
|
||||||
|
async def test_token_is_persisted_before_it_reaches_the_peer(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
seen: list[CashuTransaction] = []
|
||||||
|
|
||||||
|
async def _record_then_accept(token: str) -> dict:
|
||||||
|
seen.extend(await _sent_tokens())
|
||||||
|
return {"balance": TOPUP_SATS}
|
||||||
|
|
||||||
|
peer = _peer(0.0)
|
||||||
|
peer.topup = AsyncMock(side_effect=_record_then_accept)
|
||||||
|
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
assert [tx.token for tx in seen] == ["cashu-token-1"]
|
||||||
|
assert seen[0].collected is False
|
||||||
|
assert (await _sent_tokens())[0].collected is True
|
||||||
|
|
||||||
|
|
||||||
|
async def test_untracked_token_is_returned_and_never_sent(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
peer = _peer(0.0)
|
||||||
|
|
||||||
|
with (
|
||||||
|
_patch_wallet(auto_topup_module, peer, "cashu-token-1"),
|
||||||
|
patch.object(
|
||||||
|
auto_topup_module,
|
||||||
|
"_persist_routstr_token_and_mark_sent",
|
||||||
|
AsyncMock(side_effect=RuntimeError("database unavailable")),
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
auto_topup_module, "release_token_reservation", AsyncMock()
|
||||||
|
) as reclaim,
|
||||||
|
):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
reclaim.assert_awaited_once_with("cashu-token-1")
|
||||||
|
peer.topup.assert_not_awaited()
|
||||||
|
# Nothing left the wallet, so the slot must be free again immediately.
|
||||||
|
state = await _claim_state()
|
||||||
|
assert state is not None and state.swept is True
|
||||||
|
|
||||||
|
|
||||||
|
async def test_failed_topup_keeps_token_uncollected_and_suppresses_new_claims(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
peer = _peer(0.0, topup={"error": "rejected"})
|
||||||
|
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
tokens = await _sent_tokens()
|
||||||
|
assert len(tokens) == 1
|
||||||
|
assert tokens[0].collected is False
|
||||||
|
|
||||||
|
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
|
||||||
|
assert claim is not None and claim.phase == ROUTSTR_PHASE_SENT
|
||||||
|
|
||||||
|
peer.topup.reset_mock()
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-2"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
peer.topup.assert_not_awaited()
|
||||||
|
assert len(await _sent_tokens()) == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def test_claim_blocks_a_restarted_process(patched_db_engine: Any) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
peer = _peer(0.0, topup={"error": "rejected"})
|
||||||
|
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
# A fresh module drops every process-local variable; only a durable row
|
||||||
|
# can still stop the next payment.
|
||||||
|
reloaded = importlib.reload(auto_topup_module)
|
||||||
|
try:
|
||||||
|
peer.topup.reset_mock()
|
||||||
|
with _patch_wallet(reloaded, peer, "cashu-token-2"):
|
||||||
|
await reloaded._check_and_topup(row)
|
||||||
|
peer.topup.assert_not_awaited()
|
||||||
|
finally:
|
||||||
|
importlib.reload(auto_topup_module)
|
||||||
|
|
||||||
|
assert len(await _sent_tokens()) == 1
|
||||||
|
|
||||||
|
|
||||||
|
async def _uncredited_attempt(
|
||||||
|
row: UpstreamProviderRow, peer: MagicMock, token: str
|
||||||
|
) -> None:
|
||||||
|
"""One full payment attempt against a peer that never credits it."""
|
||||||
|
with _patch_wallet(auto_topup_module, peer, token):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
await _expire_claim()
|
||||||
|
with _patch_wallet(auto_topup_module, peer, f"{token}-retry"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
await _expire_claim()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_non_crediting_peer_is_halted_after_repeated_failures(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
peer = _peer(0.0)
|
||||||
|
|
||||||
|
for attempt in range(ROUTSTR_MAX_TOPUP_FAILURES):
|
||||||
|
await _uncredited_attempt(row, peer, f"cashu-token-{attempt}")
|
||||||
|
|
||||||
|
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
|
||||||
|
assert claim is not None and claim.phase == ROUTSTR_PHASE_HALTED
|
||||||
|
|
||||||
|
peer.topup.reset_mock()
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-after-halt"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
peer.topup.assert_not_awaited()
|
||||||
|
assert len(await _sent_tokens()) == ROUTSTR_MAX_TOPUP_FAILURES
|
||||||
|
|
||||||
|
|
||||||
|
async def _expire_claim(provider_id: int = 1) -> None:
|
||||||
|
"""Age the claim's deadline so the reconciler treats it as timed out."""
|
||||||
|
async with create_session() as session:
|
||||||
|
state = await session.get(
|
||||||
|
CashuTransaction, _routstr_state_id_for_provider(provider_id)
|
||||||
|
)
|
||||||
|
assert state is not None and state.request_id is not None
|
||||||
|
claim = _parse_routstr_request_id(state.request_id)
|
||||||
|
assert claim is not None
|
||||||
|
state.request_id = auto_topup_module._routstr_request_id(
|
||||||
|
claim.operation_id,
|
||||||
|
int(time.time()) - 1,
|
||||||
|
claim.phase,
|
||||||
|
claim.expected_sats,
|
||||||
|
claim.failures,
|
||||||
|
)
|
||||||
|
session.add(state)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_crediting_peer_releases_the_claim_for_a_later_topup(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
peer = _peer(0.0)
|
||||||
|
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
tokens = await _sent_tokens()
|
||||||
|
assert len(tokens) == 1 and tokens[0].collected is True
|
||||||
|
|
||||||
|
peer.get_balance = AsyncMock(return_value=float(TOPUP_SATS))
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-2"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
state = await _claim_state()
|
||||||
|
assert state is not None and state.collected is True
|
||||||
|
|
||||||
|
|
||||||
|
async def test_rolling_budget_refuses_a_topup_that_would_exceed_it(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
async with create_session() as session:
|
||||||
|
session.add(
|
||||||
|
CashuTransaction(
|
||||||
|
id="prior-spend",
|
||||||
|
token="cashu-prior",
|
||||||
|
amount=ROUTSTR_MAX_DAILY_TOPUP_SATS,
|
||||||
|
unit="sat",
|
||||||
|
type="out",
|
||||||
|
source="auto_topup",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
assert await _routstr_spent_last_24h_sats() == ROUTSTR_MAX_DAILY_TOPUP_SATS
|
||||||
|
|
||||||
|
peer = _peer(0.0)
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
peer.topup.assert_not_awaited()
|
||||||
|
assert await _claim_state() is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_admin_release_is_fenced_on_the_state_it_reviewed(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
peer = _peer(0.0, topup={"error": "rejected"})
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
|
||||||
|
state = await get_routstr_auto_topup_state(1)
|
||||||
|
assert state["active"] is True
|
||||||
|
assert state["phase"] == ROUTSTR_PHASE_SENT
|
||||||
|
|
||||||
|
stale = await release_routstr_auto_topup_state(
|
||||||
|
1, state_token="routstr:other:0:sent:0:0"
|
||||||
|
)
|
||||||
|
assert stale.released is False and stale.reason == "stale_state"
|
||||||
|
|
||||||
|
released = await release_routstr_auto_topup_state(
|
||||||
|
1, state_token=str(state["state_token"])
|
||||||
|
)
|
||||||
|
assert released.released is True
|
||||||
|
|
||||||
|
peer.topup.reset_mock()
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-2"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
peer.topup.assert_awaited_once()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_backoff_suppresses_retries_until_its_deadline(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed_provider()
|
||||||
|
peer = _peer(0.0)
|
||||||
|
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
await _expire_claim()
|
||||||
|
|
||||||
|
# First reconciliation after the lease: the peer never credited, so the
|
||||||
|
# claim moves to backoff rather than paying again immediately.
|
||||||
|
peer.topup.reset_mock()
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-2"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
peer.topup.assert_not_awaited()
|
||||||
|
|
||||||
|
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
|
||||||
|
assert claim is not None and claim.phase == ROUTSTR_PHASE_BACKOFF
|
||||||
|
assert claim.deadline > int(time.time())
|
||||||
@@ -0,0 +1,164 @@
|
|||||||
|
"""The Routstr top-up threshold has to state its unit.
|
||||||
|
|
||||||
|
``topup_amount_limit`` was already plain sats — it goes straight to
|
||||||
|
``send_token(amount, "sat", ...)`` and is added to the balance for logging — but
|
||||||
|
its sibling ``topup_threshold`` was compared as ``balance >= threshold * 1000``.
|
||||||
|
Two keys from one settings blob, two different units, and nothing naming either.
|
||||||
|
An operator asking for "top up below 1000 sats" got one below 1,000,000.
|
||||||
|
|
||||||
|
``topup_threshold_sats`` says what it means. Legacy values keep their effective
|
||||||
|
trigger point: reinterpreting them as sats would silently drop the trigger a
|
||||||
|
thousandfold and let a provider run dry.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from typing import Any
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from routstr.core.db import UpstreamProviderRow, create_session
|
||||||
|
from routstr.upstream import auto_topup as auto_topup_module
|
||||||
|
from routstr.upstream.auto_topup import (
|
||||||
|
_check_and_topup,
|
||||||
|
validate_routstr_auto_topup_settings,
|
||||||
|
)
|
||||||
|
|
||||||
|
from .test_routstr_auto_topup_claim import _patch_wallet, _peer, _sent_tokens
|
||||||
|
|
||||||
|
# No module-level asyncio mark: the settings cases are sync, and the suite
|
||||||
|
# already runs asyncio in auto mode.
|
||||||
|
|
||||||
|
TOPUP_SATS = 50
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _forget_legacy_hints() -> Any:
|
||||||
|
# The hint fires once per provider for the life of the process. Reach it
|
||||||
|
# through the module: a sibling test reloads auto_topup, which rebinds the
|
||||||
|
# set, so a name imported here would go on clearing the old one.
|
||||||
|
auto_topup_module._legacy_threshold_hinted.clear()
|
||||||
|
yield
|
||||||
|
auto_topup_module._legacy_threshold_hinted.clear()
|
||||||
|
|
||||||
|
|
||||||
|
async def _seed(**topup_settings: Any) -> UpstreamProviderRow:
|
||||||
|
row = UpstreamProviderRow(
|
||||||
|
id=1,
|
||||||
|
slug="peer-1",
|
||||||
|
provider_type="routstr",
|
||||||
|
base_url="https://peer.test",
|
||||||
|
api_key="secret",
|
||||||
|
enabled=True,
|
||||||
|
provider_settings=json.dumps(
|
||||||
|
{
|
||||||
|
"auto_topup": True,
|
||||||
|
"topup_amount_limit": TOPUP_SATS,
|
||||||
|
"topup_mint_url": "https://mint.test",
|
||||||
|
**topup_settings,
|
||||||
|
}
|
||||||
|
),
|
||||||
|
)
|
||||||
|
async with create_session() as session:
|
||||||
|
session.add(row)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(row)
|
||||||
|
return row
|
||||||
|
|
||||||
|
|
||||||
|
async def _topped_up(row: UpstreamProviderRow, balance: float) -> bool:
|
||||||
|
peer = _peer(balance)
|
||||||
|
with _patch_wallet(auto_topup_module, peer, "cashu-token-1"):
|
||||||
|
await _check_and_topup(row)
|
||||||
|
return bool(await _sent_tokens())
|
||||||
|
|
||||||
|
|
||||||
|
async def test_explicit_sats_threshold_is_compared_against_a_sats_balance(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed(topup_threshold_sats=1000)
|
||||||
|
assert await _topped_up(row, 1500.0) is False
|
||||||
|
|
||||||
|
|
||||||
|
async def test_explicit_sats_threshold_tops_up_below_the_stated_amount(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed(topup_threshold_sats=1000)
|
||||||
|
assert await _topped_up(row, 500.0) is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("balance", "expected"),
|
||||||
|
[(1500.0, False), (500.0, True)],
|
||||||
|
)
|
||||||
|
async def test_legacy_threshold_keeps_its_effective_trigger_point(
|
||||||
|
patched_db_engine: Any, balance: float, expected: bool
|
||||||
|
) -> None:
|
||||||
|
# 1 x 1000 == the 1000 sats the old comparison actually used. Reading it as
|
||||||
|
# 1 sat instead would leave the peer to run dry.
|
||||||
|
row = await _seed(topup_threshold=1)
|
||||||
|
assert await _topped_up(row, balance) is expected
|
||||||
|
|
||||||
|
|
||||||
|
async def test_explicit_sats_threshold_overrides_the_legacy_key(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed(topup_threshold=1, topup_threshold_sats=100)
|
||||||
|
assert await _topped_up(row, 500.0) is False
|
||||||
|
|
||||||
|
|
||||||
|
async def test_legacy_threshold_reports_the_sats_value_to_migrate_to(
|
||||||
|
patched_db_engine: Any,
|
||||||
|
) -> None:
|
||||||
|
row = await _seed(topup_threshold=1)
|
||||||
|
# The app logger does not propagate to root, so caplog would see nothing.
|
||||||
|
with patch.object(auto_topup_module, "logger") as log:
|
||||||
|
await _topped_up(row, 5000.0)
|
||||||
|
hints = [
|
||||||
|
c for c in log.warning.call_args_list if "topup_threshold_sats" in c.args[0]
|
||||||
|
]
|
||||||
|
assert len(hints) == 1
|
||||||
|
assert hints[0].kwargs["extra"]["threshold_sats"] == 1000.0
|
||||||
|
|
||||||
|
# The scheduler runs every minute; the hint must not run with it.
|
||||||
|
await _topped_up(row, 5000.0)
|
||||||
|
assert (
|
||||||
|
sum("topup_threshold_sats" in c.args[0] for c in log.warning.call_args_list)
|
||||||
|
== 1
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_settings_accept_the_sats_threshold_on_its_own() -> None:
|
||||||
|
assert (
|
||||||
|
validate_routstr_auto_topup_settings(
|
||||||
|
{
|
||||||
|
"auto_topup": True,
|
||||||
|
"topup_threshold_sats": 1000,
|
||||||
|
"topup_amount_limit": TOPUP_SATS,
|
||||||
|
"topup_mint_url": "https://mint.test",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
is None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("value", [0, -1, True, float("inf"), "1000", None])
|
||||||
|
def test_settings_reject_an_unusable_sats_threshold(value: object) -> None:
|
||||||
|
assert validate_routstr_auto_topup_settings(
|
||||||
|
{
|
||||||
|
"auto_topup": True,
|
||||||
|
"topup_threshold_sats": value,
|
||||||
|
"topup_amount_limit": TOPUP_SATS,
|
||||||
|
"topup_mint_url": "https://mint.test",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_settings_require_one_of_the_threshold_keys() -> None:
|
||||||
|
assert validate_routstr_auto_topup_settings(
|
||||||
|
{
|
||||||
|
"auto_topup": True,
|
||||||
|
"topup_amount_limit": TOPUP_SATS,
|
||||||
|
"topup_mint_url": "https://mint.test",
|
||||||
|
}
|
||||||
|
)
|
||||||
@@ -0,0 +1,255 @@
|
|||||||
|
"""The served catalog is the last guard between a stored row and a charge.
|
||||||
|
|
||||||
|
Stored pricing is JSON written by whatever produced the row — an upstream
|
||||||
|
import, an operator, a legacy migration, or a foreign writer that never passed
|
||||||
|
the admin edge. So the read path cannot assume a stored rate is a number: it
|
||||||
|
must decline to serve a row it cannot bill on, and it must survive a row it
|
||||||
|
cannot read at all rather than taking the whole catalog down with it.
|
||||||
|
|
||||||
|
The admin listing is deliberately exempt: it includes disabled models and is the
|
||||||
|
one view that still shows the operator the row that needs repair.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||||
|
|
||||||
|
from routstr.core.db import ModelRow, UpstreamProviderRow
|
||||||
|
from routstr.payment.models import list_models
|
||||||
|
|
||||||
|
_ARCHITECTURE = json.dumps(
|
||||||
|
{
|
||||||
|
"modality": "text",
|
||||||
|
"input_modalities": ["text"],
|
||||||
|
"output_modalities": ["text"],
|
||||||
|
"tokenizer": "unknown",
|
||||||
|
"instruct_type": None,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _make_provider(session: AsyncSession) -> int:
|
||||||
|
provider = UpstreamProviderRow(
|
||||||
|
provider_type="generic",
|
||||||
|
base_url="https://served-upstream.example/v1",
|
||||||
|
api_key="test-key",
|
||||||
|
provider_fee=1.0,
|
||||||
|
)
|
||||||
|
session.add(provider)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(provider)
|
||||||
|
assert provider.id is not None
|
||||||
|
return provider.id
|
||||||
|
|
||||||
|
|
||||||
|
async def _insert_row(
|
||||||
|
session: AsyncSession,
|
||||||
|
provider_id: int,
|
||||||
|
*,
|
||||||
|
model_id: str,
|
||||||
|
pricing: dict[str, object],
|
||||||
|
) -> None:
|
||||||
|
session.add(
|
||||||
|
ModelRow(
|
||||||
|
id=model_id,
|
||||||
|
name=model_id,
|
||||||
|
description="d",
|
||||||
|
created=0,
|
||||||
|
context_length=8192,
|
||||||
|
architecture=_ARCHITECTURE,
|
||||||
|
pricing=json.dumps(pricing),
|
||||||
|
upstream_provider_id=provider_id,
|
||||||
|
enabled=True,
|
||||||
|
forwarded_model_id=model_id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"bad_rate",
|
||||||
|
[float("nan"), float("inf"), -1.0],
|
||||||
|
ids=["nan", "inf", "negative"],
|
||||||
|
)
|
||||||
|
async def test_served_catalog_excludes_a_malformed_stored_rate(
|
||||||
|
integration_session: AsyncSession, bad_rate: float
|
||||||
|
) -> None:
|
||||||
|
"""A stored rate that is not a number must not be advertised.
|
||||||
|
|
||||||
|
Zero is a real price and a free model is servable, but a negative or
|
||||||
|
non-finite rate is not a price at all: serving it advertises a rate the cost
|
||||||
|
calculation cannot bill on, so every request falls through to the flat
|
||||||
|
maximum reservation — or, for a negative rate, bills an amount settlement
|
||||||
|
credits back to the caller.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="good",
|
||||||
|
pricing={"prompt": 1e-06, "completion": 2e-06},
|
||||||
|
)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="bad-rate",
|
||||||
|
pricing={"prompt": bad_rate, "completion": 2e-06},
|
||||||
|
)
|
||||||
|
|
||||||
|
served = {m.id for m in await list_models(integration_session, provider_id)}
|
||||||
|
|
||||||
|
assert served == {"good"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_free_stored_price_is_still_served(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""Zero is a real price. Rejecting malformed rates must not also drop a row
|
||||||
|
priced at zero, which is a free model and not a broken one."""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="free",
|
||||||
|
pricing={"prompt": 0.0, "completion": 0.0},
|
||||||
|
)
|
||||||
|
|
||||||
|
served = {m.id for m in await list_models(integration_session, provider_id)}
|
||||||
|
|
||||||
|
assert served == {"free"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_admin_listing_still_shows_a_malformed_stored_rate(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""The operator has to be able to see the row that needs fixing.
|
||||||
|
|
||||||
|
The backstop keeps a malformed row out of the *served* catalog. The listing
|
||||||
|
that includes disabled models is the one view where the row must still
|
||||||
|
appear, or the operator loses the ability to repair it.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="bad-rate",
|
||||||
|
pricing={"prompt": -1.0, "completion": 2e-06},
|
||||||
|
)
|
||||||
|
|
||||||
|
listed = {
|
||||||
|
m.id
|
||||||
|
for m in await list_models(
|
||||||
|
integration_session, provider_id, include_disabled=True
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
assert listed == {"bad-rate"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_one_unreadable_stored_price_does_not_blank_the_catalog(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""A single unparseable row must cost that row, not every model on the node.
|
||||||
|
|
||||||
|
Stored pricing is JSON written by whatever produced the row, so a
|
||||||
|
non-numeric rate is reachable from a legacy import or a foreign writer.
|
||||||
|
Parsing it raises out of the row-to-model conversion, and because the
|
||||||
|
conversion ran inside the catalog loop the exception took the whole listing
|
||||||
|
with it — one bad row and the node advertised nothing at all.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="good",
|
||||||
|
pricing={"prompt": 1e-06, "completion": 2e-06},
|
||||||
|
)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="unreadable",
|
||||||
|
pricing={"prompt": "not-a-number", "completion": 2e-06},
|
||||||
|
)
|
||||||
|
|
||||||
|
served = {m.id for m in await list_models(integration_session, provider_id)}
|
||||||
|
|
||||||
|
assert served == {"good"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"field",
|
||||||
|
[
|
||||||
|
"image",
|
||||||
|
"web_search",
|
||||||
|
"internal_reasoning",
|
||||||
|
"input_cache_read",
|
||||||
|
"input_cache_write",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_served_catalog_excludes_a_malformed_auxiliary_rate(
|
||||||
|
integration_session: AsyncSession, field: str
|
||||||
|
) -> None:
|
||||||
|
"""The backstop covers every billable rate, not only the token rates.
|
||||||
|
|
||||||
|
A price whose ``prompt``/``completion`` are sound can still carry a
|
||||||
|
malformed request, image, search, reasoning or cache rate — the catalog
|
||||||
|
import filter never inspects those — and the request that hits one is billed
|
||||||
|
against it just the same.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="good",
|
||||||
|
pricing={"prompt": 1e-06, "completion": 2e-06},
|
||||||
|
)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="bad-aux",
|
||||||
|
pricing={"prompt": 1e-06, "completion": 2e-06, field: -1.0},
|
||||||
|
)
|
||||||
|
|
||||||
|
served = {m.id for m in await list_models(integration_session, provider_id)}
|
||||||
|
|
||||||
|
assert served == {"good"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_a_negative_request_rate_is_clamped_on_read_and_still_served(
|
||||||
|
integration_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""``request`` is the one billable rate the row-to-model conversion repairs.
|
||||||
|
|
||||||
|
It clamps a negative stored ``request`` to zero before the price is built,
|
||||||
|
so the backstop never sees one and the row is served at a zero request rate
|
||||||
|
— money-safe, and the reason ``request`` is absent from the list of rates
|
||||||
|
above. Pinned here so that if the clamp goes, this rate joins that list
|
||||||
|
rather than quietly becoming the one unguarded field.
|
||||||
|
"""
|
||||||
|
provider_id = await _make_provider(integration_session)
|
||||||
|
await _insert_row(
|
||||||
|
integration_session,
|
||||||
|
provider_id,
|
||||||
|
model_id="neg-request",
|
||||||
|
pricing={"prompt": 1e-06, "completion": 2e-06, "request": -1.0},
|
||||||
|
)
|
||||||
|
|
||||||
|
served = await list_models(integration_session, provider_id)
|
||||||
|
|
||||||
|
assert [m.id for m in served] == ["neg-request"]
|
||||||
|
assert served[0].pricing.request == 0.0
|
||||||
@@ -1,210 +0,0 @@
|
|||||||
"""
|
|
||||||
Integration tests for reactive swap fee retries via the wallet topup endpoint.
|
|
||||||
|
|
||||||
Foreign-mint tokens are swapped to the primary mint using the foreign mint's
|
|
||||||
melt quote, whose fee_reserve is a non-binding estimate (NUT-05): the mint may
|
|
||||||
demand more when re-quoting or at melt execution. These tests cover the
|
|
||||||
endpoint behaviour in those cases:
|
|
||||||
|
|
||||||
1. The mint demands one sat more at melt time than every quote reported
|
|
||||||
(the mint.cubabitcoin.org incident): the swap retries with a smaller
|
|
||||||
invoice and the topup succeeds, crediting the recomputed amount.
|
|
||||||
2. The real melt quote reports a higher fee_reserve than the estimate: the
|
|
||||||
swap re-quotes from the observed fee and the topup succeeds.
|
|
||||||
3. The mint escalates its fee demands on every attempt: the retry budget is
|
|
||||||
exhausted and the endpoint returns 400 with a clear error (never 500),
|
|
||||||
without ever executing a melt.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from collections.abc import Callable
|
|
||||||
from unittest.mock import AsyncMock, Mock, patch
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
from cashu.core.base import MeltQuoteState
|
|
||||||
from httpx import AsyncClient, Response
|
|
||||||
|
|
||||||
from routstr.core.settings import settings
|
|
||||||
|
|
||||||
# Captured at collection time, before the integration_app fixture replaces it
|
|
||||||
# with the testmint stub that bypasses swapping (see conftest.py).
|
|
||||||
from routstr.wallet import recieve_token as _real_recieve_token
|
|
||||||
|
|
||||||
# Match the authenticated fixture's persisted refund mint: existing-key topups
|
|
||||||
# are intentionally constrained to that mint for collateral provenance.
|
|
||||||
PRIMARY_MINT = "http://localhost:3338"
|
|
||||||
|
|
||||||
|
|
||||||
def _make_swap_mocks(
|
|
||||||
token_amount: int,
|
|
||||||
fee_reserves: list[int],
|
|
||||||
input_fees: int = 0,
|
|
||||||
mint_url: str = "http://foreign-mint:3338",
|
|
||||||
) -> tuple[Mock, Mock, Mock]:
|
|
||||||
"""Return (token, token_wallet, primary_wallet) mocks that act like a mint.
|
|
||||||
|
|
||||||
Mint quotes pass the requested amount through their ``request`` field and
|
|
||||||
melt quotes echo that amount back, so the mocks stay consistent for
|
|
||||||
whatever amounts the implementation requests. ``fee_reserves`` supplies the
|
|
||||||
fee_reserve of each successive melt quote (the first serves the estimation
|
|
||||||
pass); requesting more quotes than provided fails the test.
|
|
||||||
"""
|
|
||||||
mock_token = Mock()
|
|
||||||
mock_token.mint = mint_url
|
|
||||||
mock_token.unit = "sat"
|
|
||||||
mock_token.amount = token_amount
|
|
||||||
mock_token.keysets = ["keyset1"]
|
|
||||||
mock_token.proofs = [Mock(amount=token_amount)]
|
|
||||||
|
|
||||||
mock_token_wallet = Mock()
|
|
||||||
mock_token_wallet.load_mint_keysets = AsyncMock()
|
|
||||||
mock_token_wallet.activate_keyset = AsyncMock()
|
|
||||||
mock_token_wallet._expand_short_keyset_ids = AsyncMock()
|
|
||||||
mock_token_wallet.load_proofs = AsyncMock()
|
|
||||||
mock_token_wallet.get_fees_for_proofs = Mock(return_value=input_fees)
|
|
||||||
|
|
||||||
mock_primary_wallet = Mock()
|
|
||||||
mock_primary_wallet.load_mint = AsyncMock()
|
|
||||||
mock_primary_wallet.load_proofs = AsyncMock()
|
|
||||||
mock_primary_wallet.available_balance = Mock(amount=0)
|
|
||||||
mock_primary_wallet.mint = AsyncMock(return_value=Mock())
|
|
||||||
|
|
||||||
fees = iter(fee_reserves)
|
|
||||||
|
|
||||||
def _next_fee() -> int:
|
|
||||||
try:
|
|
||||||
return next(fees)
|
|
||||||
except StopIteration:
|
|
||||||
raise AssertionError(
|
|
||||||
"more melt quotes requested than fee_reserves provided"
|
|
||||||
) from None
|
|
||||||
|
|
||||||
mock_primary_wallet.request_mint = AsyncMock(
|
|
||||||
side_effect=lambda amount: Mock(quote=f"mint_quote_{amount}", request=amount)
|
|
||||||
)
|
|
||||||
mock_token_wallet.melt_quote = AsyncMock(
|
|
||||||
side_effect=lambda invoice: Mock(
|
|
||||||
quote=f"melt_quote_{invoice}", amount=invoice, fee_reserve=_next_fee()
|
|
||||||
)
|
|
||||||
)
|
|
||||||
mock_token_wallet.melt = AsyncMock(
|
|
||||||
return_value=Mock(state=MeltQuoteState.paid)
|
|
||||||
)
|
|
||||||
|
|
||||||
return mock_token, mock_token_wallet, mock_primary_wallet
|
|
||||||
|
|
||||||
|
|
||||||
def _wallet_router(primary_wallet: Mock, token_wallet: Mock) -> Callable[..., Mock]:
|
|
||||||
"""Route get_wallet calls to the primary or foreign wallet mock by URL."""
|
|
||||||
|
|
||||||
def fake_get_wallet(
|
|
||||||
mint_url: str,
|
|
||||||
unit: str = "sat",
|
|
||||||
load: bool = True,
|
|
||||||
**kwargs: object,
|
|
||||||
) -> Mock:
|
|
||||||
return primary_wallet if mint_url == PRIMARY_MINT else token_wallet
|
|
||||||
|
|
||||||
return fake_get_wallet
|
|
||||||
|
|
||||||
|
|
||||||
async def _post_topup(
|
|
||||||
client: AsyncClient,
|
|
||||||
mock_token: Mock,
|
|
||||||
token_wallet: Mock,
|
|
||||||
primary_wallet: Mock,
|
|
||||||
) -> Response:
|
|
||||||
"""POST /v1/wallet/topup with the swap layer mocked at the mint boundary.
|
|
||||||
|
|
||||||
The conftest's testmint stub for recieve_token is swapped back for the
|
|
||||||
real implementation so the request exercises the actual swap path.
|
|
||||||
"""
|
|
||||||
with patch("routstr.wallet.recieve_token", _real_recieve_token):
|
|
||||||
with patch(
|
|
||||||
"routstr.wallet.deserialize_token_from_string", return_value=mock_token
|
|
||||||
):
|
|
||||||
with patch(
|
|
||||||
"routstr.wallet.get_wallet",
|
|
||||||
side_effect=_wallet_router(primary_wallet, token_wallet),
|
|
||||||
):
|
|
||||||
with patch.object(settings, "primary_mint", PRIMARY_MINT):
|
|
||||||
with patch.object(settings, "primary_mint_unit", "sat"):
|
|
||||||
with patch.object(settings, "cashu_mints", [PRIMARY_MINT]):
|
|
||||||
return await client.post(
|
|
||||||
"/v1/wallet/topup",
|
|
||||||
params={"cashu_token": "cashuAtest_foreign_token"},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_topup_retries_when_melt_demands_more_than_quoted(
|
|
||||||
authenticated_client: AsyncClient,
|
|
||||||
) -> None:
|
|
||||||
"""A 179-sat token where every quote reports fee_reserve=1 but the mint
|
|
||||||
rejects the first melt demanding 180. The retry shrinks the invoice to 177
|
|
||||||
and the topup credits 177 sats (177_000 msats)."""
|
|
||||||
mock_token, token_wallet, primary_wallet = _make_swap_mocks(
|
|
||||||
179, fee_reserves=[1, 1, 1], mint_url="http://mint.cubabitcoin.org"
|
|
||||||
)
|
|
||||||
token_wallet.melt.side_effect = [
|
|
||||||
Exception(
|
|
||||||
"Mint Error: not enough inputs provided for melt. "
|
|
||||||
"Provided: 179, needed: 180 (Code: 11000)"
|
|
||||||
),
|
|
||||||
Mock(state=MeltQuoteState.paid),
|
|
||||||
]
|
|
||||||
|
|
||||||
response = await _post_topup(
|
|
||||||
authenticated_client, mock_token, token_wallet, primary_wallet
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert response.json()["msats"] == 177_000
|
|
||||||
assert token_wallet.melt.call_count == 2
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_topup_retries_when_quote_fee_exceeds_estimate(
|
|
||||||
authenticated_client: AsyncClient,
|
|
||||||
) -> None:
|
|
||||||
"""A 1000-sat token estimated at fee 20, but the real quote demands 23.
|
|
||||||
The retry recomputes 1000 - 23 = 977, which fits, and the topup credits
|
|
||||||
977 sats (977_000 msats) with a single melt."""
|
|
||||||
mock_token, token_wallet, primary_wallet = _make_swap_mocks(
|
|
||||||
1000, fee_reserves=[20, 23, 23]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await _post_topup(
|
|
||||||
authenticated_client, mock_token, token_wallet, primary_wallet
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert response.json()["msats"] == 977_000
|
|
||||||
assert token_wallet.melt.call_count == 1
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_topup_returns_422_when_retries_exhausted(
|
|
||||||
authenticated_client: AsyncClient,
|
|
||||||
) -> None:
|
|
||||||
"""A mint that escalates fee_reserve on every re-quote (1 → 10 → 25 → 50)
|
|
||||||
exhausts the retry budget: clean 422 mint_error/too-small taxonomy, melt
|
|
||||||
never executed."""
|
|
||||||
mock_token, token_wallet, primary_wallet = _make_swap_mocks(
|
|
||||||
1000, fee_reserves=[1, 10, 25, 50]
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await _post_topup(
|
|
||||||
authenticated_client, mock_token, token_wallet, primary_wallet
|
|
||||||
)
|
|
||||||
|
|
||||||
assert response.status_code == 422
|
|
||||||
raw_detail = response.json()["detail"]
|
|
||||||
message = (
|
|
||||||
raw_detail["error"]["message"] if isinstance(raw_detail, dict) else raw_detail
|
|
||||||
)
|
|
||||||
assert "too small to cover swap fees" in message
|
|
||||||
assert token_wallet.melt_quote.call_count == 4 # estimation + 3 attempts
|
|
||||||
token_wallet.melt.assert_not_called()
|
|
||||||
@@ -22,18 +22,18 @@ async def _add_key(
|
|||||||
hashed_key: str,
|
hashed_key: str,
|
||||||
*,
|
*,
|
||||||
balance: int = 0,
|
balance: int = 0,
|
||||||
|
reserved_balance: int = 0,
|
||||||
total_spent: int = 0,
|
total_spent: int = 0,
|
||||||
total_requests: int = 0,
|
total_requests: int = 0,
|
||||||
created_at: int | None = None,
|
created_at: int | None = None,
|
||||||
parent_key_hash: str | None = None,
|
|
||||||
refund_address: str | None = None,
|
refund_address: str | None = None,
|
||||||
) -> ApiKey:
|
) -> ApiKey:
|
||||||
key = ApiKey(
|
key = ApiKey(
|
||||||
hashed_key=hashed_key,
|
hashed_key=hashed_key,
|
||||||
balance=balance,
|
balance=balance,
|
||||||
|
reserved_balance=reserved_balance,
|
||||||
total_spent=total_spent,
|
total_spent=total_spent,
|
||||||
total_requests=total_requests,
|
total_requests=total_requests,
|
||||||
parent_key_hash=parent_key_hash,
|
|
||||||
refund_address=refund_address,
|
refund_address=refund_address,
|
||||||
)
|
)
|
||||||
key.created_at = created_at
|
key.created_at = created_at
|
||||||
@@ -123,29 +123,18 @@ async def test_temporary_balances_pagination(
|
|||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_temporary_balances_totals_exclude_child_balance(
|
async def test_temporary_balances_totals(
|
||||||
integration_client: httpx.AsyncClient,
|
integration_client: httpx.AsyncClient,
|
||||||
integration_session: AsyncSession,
|
integration_session: AsyncSession,
|
||||||
) -> None:
|
) -> None:
|
||||||
await _add_key(
|
await _add_key(
|
||||||
integration_session,
|
integration_session,
|
||||||
"parent",
|
"standalone_key",
|
||||||
balance=5000,
|
balance=5000,
|
||||||
total_spent=100,
|
total_spent=100,
|
||||||
total_requests=3,
|
total_requests=3,
|
||||||
created_at=1000,
|
created_at=1000,
|
||||||
)
|
)
|
||||||
# Child draws from parent's balance, so its balance must NOT be summed,
|
|
||||||
# but its spent/requests still count.
|
|
||||||
await _add_key(
|
|
||||||
integration_session,
|
|
||||||
"child",
|
|
||||||
balance=0,
|
|
||||||
total_spent=200,
|
|
||||||
total_requests=7,
|
|
||||||
created_at=1001,
|
|
||||||
parent_key_hash="parent",
|
|
||||||
)
|
|
||||||
|
|
||||||
response = await integration_client.get(
|
response = await integration_client.get(
|
||||||
"/admin/api/temporary-balances", headers=_admin_headers()
|
"/admin/api/temporary-balances", headers=_admin_headers()
|
||||||
@@ -153,8 +142,8 @@ async def test_temporary_balances_totals_exclude_child_balance(
|
|||||||
|
|
||||||
totals = response.json()["totals"]
|
totals = response.json()["totals"]
|
||||||
assert totals["total_balance"] == 5000
|
assert totals["total_balance"] == 5000
|
||||||
assert totals["total_spent"] == 300
|
assert totals["total_spent"] == 100
|
||||||
assert totals["total_requests"] == 10
|
assert totals["total_requests"] == 3
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user