Compare commits

...

97 Commits

Author SHA1 Message Date
9qeklajc
58b468018a update migration 2026-04-28 00:42:12 +02:00
root
2b79831c61 Merge branch 'main' into dynamic-prices 2026-04-28 00:33:13 +02:00
9qeklajc
cd7de3958c Merge pull request #478 from Routstr/missing-model-pricing
make sure to always emit sats cost
2026-04-28 00:32:32 +02:00
9qeklajc
8c0ac499ef make sure to always emit sats cost 2026-04-28 00:01:27 +02:00
9qeklajc
28d91227af Merge pull request #475 from Routstr/admin-token
Admin token
2026-04-26 22:32:21 +02:00
9qeklajc
a2db3e2d57 fmt 2026-04-26 22:19:30 +02:00
9qeklajc
e43ceb2e43 Merge pull request #477 from Routstr/model-refresh
enforce-models-refresh-from-upstream
2026-04-26 22:03:59 +02:00
9qeklajc
689a07f562 enforce-models-refresh-from-upstream 2026-04-26 21:41:38 +02:00
9qeklajc
da487a850e Merge branch 'main' into admin-token 2026-04-26 00:09:00 +02:00
9qeklajc
fa6d3c76d0 Merge pull request #474 from Routstr/fix-field-label
use correct field label
2026-04-26 00:05:13 +02:00
9qeklajc
ede1804d4b use correct field label 2026-04-25 23:57:24 +02:00
9qeklajc
3392e8d4cb Merge pull request #473 from Routstr/bump-release-version
release v0.4.3
2026-04-25 11:56:48 +02:00
9qeklajc
b5174d9753 release v0.4.3 2026-04-25 11:54:54 +02:00
9qeklajc
1f2ff8a99c added admin token 2026-04-25 11:18:58 +02:00
9qeklajc
0c60644ba2 Merge pull request #472 from Routstr/fix-revision
fix migration
2026-04-24 23:08:06 +02:00
9qeklajc
aca8d43a61 fix migration 2026-04-24 23:05:09 +02:00
9qeklajc
7fe4c1963b Merge pull request #470 from Routstr/add-dev-cut
update migration rev id
2026-04-24 15:21:32 +02:00
9qeklajc
9a0919f149 update migration rev id 2026-04-24 15:19:53 +02:00
9qeklajc
d69ab913d4 Merge pull request #419 from Routstr/add-dev-cut
add-dev-cut
2026-04-23 23:22:23 +02:00
9qeklajc
724de338f2 update migration 2026-04-23 00:02:42 +02:00
9qeklajc
c98cc30fd4 Merge branch 'main' into dynamic-prices 2026-04-22 23:57:55 +02:00
9qeklajc
42b8c332df Merge pull request #467 from Routstr/fix-concurent-refund-with-adding-token-history
Fix concurent refund with adding token history
2026-04-22 23:45:38 +02:00
9qeklajc
aa682bf8ec fix test 2026-04-22 23:43:11 +02:00
9qeklajc
c53e72e80a update migration 2026-04-22 23:31:49 +02:00
9qeklajc
afd81aeca2 reduce retries 2026-04-22 23:14:45 +02:00
9qeklajc
4e03145323 clean up 2026-04-22 23:13:29 +02:00
9qeklajc
169686681f Merge branch 'main' into add-dev-cut 2026-04-22 23:03:35 +02:00
9qeklajc
fa7d2804bb add test 2026-04-22 23:00:04 +02:00
9qeklajc
ccabc5d06d Merge branch 'main' into fix-response-error-forwarding 2026-04-22 22:17:59 +02:00
9qeklajc
cb68227c88 fix race cond. test 2026-04-22 22:17:41 +02:00
9qeklajc
c72fc7dd56 Merge pull request #466 from Routstr/add-detailed-response
fix forwarding upstream error responses
2026-04-22 22:16:41 +02:00
9qeklajc
8c6d1f89dc fix forwarding upstream error responses 2026-04-22 22:12:44 +02:00
9qeklajc
1ab7e54bd7 improve history collection 2026-04-22 22:09:13 +02:00
9qeklajc
42dceb0cd6 make table scrollable 2026-04-22 21:57:29 +02:00
9qeklajc
05115c3387 add api-key history and fix race condition while topup 2026-04-22 21:50:22 +02:00
9qeklajc
c0c8cafd00 fix forwarding upstream error responses 2026-04-20 16:45:05 +02:00
root
4691294b24 Merge branch 'main' into dynamic-prices 2026-04-20 13:10:00 +02:00
9qeklajc
16d6d66d17 Merge pull request #460 from Routstr/fix-test-interface-ui
Fix test interface UI
2026-04-19 15:29:00 +02:00
9qeklajc
3681fc8aab no advanced testing for now 2026-04-19 15:22:47 +02:00
9qeklajc
d95c09e00c fix test model connection 2026-04-19 15:21:23 +02:00
9qeklajc
3cdef5ec17 Merge pull request #459 from Routstr/add-detailed-logging
add logging reason for bad not succeed requests
2026-04-18 00:33:39 +02:00
9qeklajc
58e1620347 add logging reason for bad not succeed requests 2026-04-18 00:30:19 +02:00
9qeklajc
42b5a5e6c8 Merge pull request #458 from Routstr/add-cost-usage-to-messages-endpoint
Add cost usage to messages endpoint
2026-04-16 17:36:18 +02:00
9qeklajc
d84d249f2c Merge branch 'main' into add-dev-cut
# Conflicts:
#	routstr/auth.py
2026-04-15 23:07:37 +02:00
9qeklajc
8b6393f794 clean up 2026-04-15 21:56:30 +02:00
9qeklajc
111060c004 Merge branch 'add-cost-usage-to-messages-endpoint' into dynamic-prices 2026-04-15 21:32:38 +02:00
9qeklajc
81146710ba revert 2026-04-15 21:32:33 +02:00
9qeklajc
f2b700a2e6 Merge branch 'add-cost-usage-to-messages-endpoint' into dynamic-prices 2026-04-15 21:27:09 +02:00
9qeklajc
9008b00fc9 revert 2026-04-15 21:27:02 +02:00
9qeklajc
d7467dbad2 Merge branch 'add-cost-usage-to-messages-endpoint' into dynamic-prices 2026-04-15 21:22:38 +02:00
9qeklajc
25f39897d7 mirror logic from chat completion 2026-04-15 20:45:13 +02:00
9qeklajc
45bdfaee58 best effort to be compatible with legacy code 2026-04-15 20:18:56 +02:00
9qeklajc
a1ae6e94e9 match model when versioned 2026-04-15 20:05:36 +02:00
9qeklajc
1190a09acb Merge branch 'add-cost-usage-to-messages-endpoint' into dynamic-prices 2026-04-15 19:20:56 +02:00
9qeklajc
29f52116c3 match model when versioned 2026-04-15 19:20:50 +02:00
9qeklajc
eebcc67c85 add cost usages 2026-04-15 17:19:13 +02:00
9qeklajc
3f9e7f7728 Merge pull request #457 from Routstr/fix-ppq-model-fetching
fix ppq model fetching
2026-04-14 18:49:26 +02:00
9qeklajc
a5b5549edd fix ppq model fetching 2026-04-14 18:46:38 +02:00
9qeklajc
85532f32a1 Merge branch 'main' into dynamic-prices 2026-04-14 00:46:00 +02:00
9qeklajc
22eec0162c Merge pull request #456 from Routstr/fix-serializing-broken-chunks
make streaming serialization more robust
2026-04-14 00:45:48 +02:00
9qeklajc
9f55da9bb8 make streaming serialization more robust 2026-04-14 00:44:10 +02:00
9qeklajc
ea8257c44b Merge branch 'main' into dynamic-prices 2026-04-14 00:32:22 +02:00
9qeklajc
16dce9ea81 Merge pull request #455 from Routstr/check-primary-mint
enforce no swap when token from primary mint
2026-04-14 00:32:06 +02:00
9qeklajc
963ee04619 enforce no swap when token from primary mint 2026-04-14 00:29:54 +02:00
9qeklajc
d89e740ff8 Merge branch 'main' into dynamic-prices
# Conflicts:
#	routstr/upstream/base.py
2026-04-13 23:52:46 +02:00
9qeklajc
b874b1f01c Merge pull request #454 from Routstr/response-id-should-be-set
upstream drops id field which breaks opencode flow
2026-04-13 23:27:41 +02:00
9qeklajc
b6cca3d3a0 upstream drops id field which breaks opencode flow 2026-04-13 23:11:07 +02:00
9qeklajc
86ebc84f4c Merge pull request #453 from Routstr/remove-hardcoded-fee
remove hardcoded deduction (legacy code)
2026-04-13 22:47:23 +02:00
9qeklajc
8a74c0543f remove hardcoded deduction (legacy code) 2026-04-13 22:41:04 +02:00
9qeklajc
af4be5bfec Merge branch 'main' into dynamic-prices 2026-04-12 22:46:18 +02:00
9qeklajc
3d16a9b988 Merge pull request #449 from Routstr/apikey-transactions
add apikey refund tracking
2026-04-12 22:38:40 +02:00
9qeklajc
40f98b99aa set collected state 2026-04-12 22:36:52 +02:00
9qeklajc
df2a925577 added-dynamic-price-setting 2026-04-12 17:38:35 +02:00
9qeklajc
7fcae5b08d Merge pull request #451 from Routstr/feature/update-deployment-docs
Updated deployment docs
2026-04-12 11:23:10 +02:00
redshift
f50cb31749 Updated deployment docs 2026-04-11 22:15:50 +01:00
9qeklajc
85a1ea4b4c add apikey refund tracking 2026-04-10 15:34:19 +02:00
9qeklajc
0f0f8c40bf Merge pull request #447 from Routstr/436-update-model-id
add support to using custom model ids
2026-04-10 14:18:24 +02:00
9qeklajc
a40d224ee7 Merge pull request #446 from Routstr/improve-negative-balance-handling
Improve negative balance handling
2026-04-10 14:14:22 +02:00
9qeklajc
a637acd8f4 use db fixture 2026-04-09 23:08:01 +02:00
9qeklajc
58be0c7976 clean u 2026-04-09 20:42:36 +02:00
9qeklajc
d934f3eead clean up 2026-04-09 20:41:59 +02:00
9qeklajc
74b58e5fa3 add some tests and improve payment handling 2026-04-09 20:25:43 +02:00
9qeklajc
c1497e0cfe Merge pull request #444 from Routstr/add-model-id-to-validation-log
add model id to validation log
2026-04-09 01:13:07 +02:00
9qeklajc
30db582321 add model id to validation log 2026-04-08 00:36:05 +02:00
9qeklajc
a1018776e9 make sure no negative balance can happen 2026-04-08 00:16:55 +02:00
9qeklajc
5ef499e2ae Merge pull request #443 from Routstr/437-clean-up-logging
do not log client host address
2026-04-07 00:44:07 +02:00
9qeklajc
c7c802c610 do not log client host address 2026-04-07 00:23:53 +02:00
9qeklajc
685368bb0a add support to using custom model ids 2026-04-06 20:09:12 +02:00
9qeklajc
55e240d92a Merge pull request #435 from Routstr/add-missing-sat-cost-in-x-cashu
add sat cost to x-cashu response
2026-04-05 01:35:17 +02:00
9qeklajc
7da4ad3818 add sat cost to x-cashu response 2026-04-05 01:20:58 +02:00
9qeklajc
453337cb2c update default payout 2026-04-05 00:36:59 +02:00
9qeklajc
9e1934bfda Merge pull request #434 from Routstr/opencode-display-usage-correctly
Opencode display usage correctly
2026-04-05 00:34:22 +02:00
9qeklajc
d686e0e851 Merge pull request #433 from Routstr/update-x-cashu-swept-default-to-one-week
x-cashu swept default to one week
2026-04-05 00:23:37 +02:00
9qeklajc
7708ed1c8b x-cashu swept default to one week 2026-04-05 00:21:53 +02:00
9qeklajc
236854bfe4 update lightining address 2026-03-25 10:25:46 +01:00
9qeklajc
a7886c528f Merge branch 'main' into add-dev-cut 2026-03-25 10:24:55 +01:00
9qeklajc
a7b815b29f add-dev-cut 2026-03-23 20:07:24 +01:00
58 changed files with 7579 additions and 2735 deletions

View File

@@ -51,14 +51,26 @@ curl https://api.routstr.com/v1/chat/completions \
## Quick Start (Docker)
If you are a node runner, start a Routstr Core instance and configure upstream access in the dashboard.
If you are a node runner, start a Routstr Core instance using Docker Compose:
```bash
docker run -d \
--name routstr-proxy \
-p 8000:8000 \
ghcr.io/routstr/proxy:latest
```
1. **Prepare your `.env`**:
```bash
ADMIN_PASSWORD=mysecretpassword
NAME="My AI Node"
DESCRIPTION="Fast access to models"
NSEC=yournsec
RECEIVE_LN_ADDRESS=yourname@wallet.com
```
2. **Start the services**:
```bash
docker compose up -d
```
3. **Configure**:
Open [http://localhost:8000/admin/](http://localhost:8000/admin/) to connect your AI providers and set pricing.
For full instructions, see the **[Provider Quick Start Guide](https://docs.routstr.com/provider/quickstart/)**.
## Development

View File

@@ -6,16 +6,7 @@ Production deployment guide for Routstr Provider nodes.
For production, use Docker Compose with persistent storage and optional Tor support.
### Unified Setup (All-in-one)
To build and run the node with the UI integrated in a single container using the multi-stage build:
```bash
docker build -f Dockerfile.full -t routstr-full .
docker run -d -p 8000:8000 --env-file .env routstr-full
```
### Advanced Setup (Separated UI & Node)
Use the included `compose.yml` for a more flexible setup that separates the UI build process from the node execution. This is useful for development or when you want to manage Tor as a separate service.
Use the included `compose.yml` for a flexible setup that handles both the UI and the node execution. This is useful for development or when you want to manage Tor as a separate service.
```bash
docker compose up -d
@@ -184,20 +175,16 @@ docker compose up -d
## Building from Source
### Unified Image (UI + Node)
The easiest way to build everything from source into a single production-ready image:
### Using Docker Compose
The easiest way to build everything from source:
```bash
docker build -f Dockerfile.full -t routstr-full .
docker compose build
```
### Individual Components
If you prefer building them separately or using Docker Compose:
If you prefer building the node only (requires manual UI build first):
```bash
# Build using compose
docker compose build
# Or build the node only (requires manual UI build first)
docker build -t routstr-node .
```

View File

@@ -35,6 +35,7 @@ ADMIN_PASSWORD=mysecretpassword
# Node Identity
NAME="My AI Node"
DESCRIPTION="Fast access to models"
NSEC=yournsec
# Lightning Payouts
RECEIVE_LN_ADDRESS=yourname@wallet.com
@@ -43,32 +44,10 @@ RECEIVE_LN_ADDRESS=yourname@wallet.com
## 2. Start the Node
You can run the pre-built image directly:
The recommended way to run Routstr is using Docker Compose, which handles the node, the UI, and optional services like Tor.
```bash
docker run -d \
--name routstr \
-p 8000:8000 \
--env-file .env \
-v routstr-data:/app/data \
ghcr.io/routstr/proxy:latest
```
*Note: The pre-built image does not contain the UI. For the all-in-one experience with the Admin Dashboard, use the Build from Source instructions below.*
### Build from Source (Recommended)
If you want to build the node and UI yourself from source, use the unified Dockerfile:
```bash
git clone https://github.com/routstr/routstr-core.git
cd routstr-core
# Edit your .env with ADMIN_PASSWORD and API keys
cp .env.example .env
nano .env
docker build -f Dockerfile.full -t routstr-local .
docker run -d -p 8000:8000 --env-file .env --name routstr routstr-local
docker compose up -d
```
Verify it's running:
@@ -77,6 +56,15 @@ Verify it's running:
curl http://localhost:8000/v1/info
```
### Build from Source (Optional)
If you've cloned the repository and want to build the images yourself:
```bash
docker compose build
docker compose up -d
```
---
## 3. Configure via Dashboard

View File

@@ -0,0 +1,32 @@
"""add routstr_fees table
Revision ID: 02650cd6f028
Revises: c3d4e5f6a7b8
Create Date: 2026-04-24 00:00:00.000000
"""
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "02650cd6f028"
down_revision = "c3d4e5f6a7b8"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"routstr_fees",
sa.Column("id", sa.Integer(), nullable=False),
sa.Column("accumulated_msats", sa.Integer(), nullable=False, server_default="0"),
sa.Column("total_paid_msats", sa.Integer(), nullable=False, server_default="0"),
sa.Column("last_paid_at", sa.Integer(), nullable=True),
sa.PrimaryKeyConstraint("id"),
)
# Seed with a single row
op.execute("INSERT INTO routstr_fees (id, accumulated_msats, total_paid_msats) VALUES (1, 0, 0)")
def downgrade() -> None:
op.drop_table("routstr_fees")

View File

@@ -0,0 +1,57 @@
"""add provider_fee_schedules and provider_fee_default to upstream_providers
Revision ID: 6d2fa295fa43
Revises: cli_tokens_001
Create Date: 2026-04-28 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "6d2fa295fa43"
down_revision = "cli_tokens_001"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {c["name"] for c in inspector.get_columns("upstream_providers")}
if "provider_fee_default" not in columns:
op.add_column(
"upstream_providers",
sa.Column(
"provider_fee_default",
sa.Float(),
nullable=False,
server_default="1.01",
),
)
# Preserve any custom per-provider fees by copying from provider_fee.
op.execute(
"UPDATE upstream_providers "
"SET provider_fee_default = provider_fee "
"WHERE provider_fee IS NOT NULL"
)
if "provider_fee_schedules" not in columns:
op.add_column(
"upstream_providers",
sa.Column("provider_fee_schedules", sa.Text(), nullable=True),
)
def downgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = {c["name"] for c in inspector.get_columns("upstream_providers")}
if "provider_fee_schedules" in columns:
op.drop_column("upstream_providers", "provider_fee_schedules")
if "provider_fee_default" in columns:
op.drop_column("upstream_providers", "provider_fee_default")

View File

@@ -0,0 +1,33 @@
"""add forwarded_model_id to models
Revision ID: b1c2d3e4f5a6
Revises: a776ca70e5fe
Create Date: 2026-04-05 00:00:00.000000
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
# revision identifiers, used by Alembic.
revision = "b1c2d3e4f5a6"
down_revision = "a776ca70e5fe"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"models",
sa.Column(
"forwarded_model_id",
sqlmodel.sql.sqltypes.AutoString(),
nullable=True,
),
)
# Backfill: set forwarded_model_id = id for all existing rows
op.execute("UPDATE models SET forwarded_model_id = id WHERE forwarded_model_id IS NULL")
def downgrade() -> None:
op.drop_column("models", "forwarded_model_id")

View File

@@ -0,0 +1,36 @@
"""add source to cashu_transactions
Revision ID: c3d4e5f6a7b8
Revises: b1c2d3e4f5a6
Create Date: 2026-04-10 00:00:00.000000
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
# revision identifiers, used by Alembic.
revision = "c3d4e5f6a7b8"
down_revision = "b1c2d3e4f5a6"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = [col["name"] for col in inspector.get_columns("cashu_transactions")]
if "source" not in columns:
op.add_column(
"cashu_transactions",
sa.Column(
"source",
sqlmodel.sql.sqltypes.AutoString(),
nullable=False,
server_default="x-cashu",
),
)
def downgrade() -> None:
op.drop_column("cashu_transactions", "source")

View File

@@ -0,0 +1,34 @@
"""add cli_tokens table
Revision ID: cli_tokens_001
Revises: e8f9a0b1c2d3
Create Date: 2026-04-25 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "cli_tokens_001"
down_revision = "e8f9a0b1c2d3"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"cli_tokens",
sa.Column("id", sa.String(), primary_key=True, nullable=False),
sa.Column("token", sa.String(), nullable=False, unique=True),
sa.Column("name", sa.String(), nullable=False),
sa.Column("created_at", sa.Integer(), nullable=False),
sa.Column("last_used_at", sa.Integer(), nullable=True),
sa.Column("expires_at", sa.Integer(), nullable=True),
)
op.create_index("ix_cli_tokens_token", "cli_tokens", ["token"], unique=True)
def downgrade() -> None:
op.drop_index("ix_cli_tokens_token", table_name="cli_tokens")
op.drop_table("cli_tokens")

View File

@@ -0,0 +1,46 @@
"""add api key link to cashu_transactions
Revision ID: d4e5f6a7b8c9
Revises: c3d4e5f6a7b8
Create Date: 2026-04-20 00:00:00.000000
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
# revision identifiers, used by Alembic.
revision = "d4e5f6a7b8c9"
down_revision = "c3d4e5f6a7b8"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
columns = [col["name"] for col in inspector.get_columns("cashu_transactions")]
indexes = {index["name"] for index in inspector.get_indexes("cashu_transactions")}
if "api_key_hashed_key" not in columns:
op.add_column(
"cashu_transactions",
sa.Column(
"api_key_hashed_key",
sqlmodel.sql.sqltypes.AutoString(),
nullable=True,
),
)
if "ix_cashu_transactions_api_key_hashed_key" not in indexes:
op.create_index(
"ix_cashu_transactions_api_key_hashed_key",
"cashu_transactions",
["api_key_hashed_key"],
unique=False,
)
def downgrade() -> None:
op.drop_index("ix_cashu_transactions_api_key_hashed_key", table_name="cashu_transactions")
op.drop_column("cashu_transactions", "api_key_hashed_key")

View File

@@ -0,0 +1,20 @@
"""merge heads: routstr_fees + api_key_to_cashu_transactions
Revision ID: e8f9a0b1c2d3
Revises: 02650cd6f028, d4e5f6a7b8c9
Create Date: 2026-04-24 00:00:00.000000
"""
# revision identifiers, used by Alembic.
revision = "e8f9a0b1c2d3"
down_revision = ("02650cd6f028", "d4e5f6a7b8c9")
branch_labels = None
depends_on = None
def upgrade() -> None:
pass
def downgrade() -> None:
pass

View File

@@ -1,6 +1,6 @@
[project]
name = "routstr"
version = "0.4.1"
version = "0.4.3"
description = "Payment proxy for your LLM endpoint using cashu and nostr."
readme = "README.md"
requires-python = ">=3.11"

View File

@@ -217,6 +217,10 @@ def create_model_mappings(
if prefixed_id not in aliases:
aliases.append(prefixed_id)
# Register forwarded_model_id as a routable alias
if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases:
aliases.append(model_to_use.forwarded_model_id)
# Try to set each alias
for alias in aliases:
_add_candidate(alias, model_to_use, upstream)
@@ -305,6 +309,10 @@ def create_model_mappings(
if prefixed_id not in aliases:
aliases.append(prefixed_id)
# Register forwarded_model_id as a routable alias
if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases:
aliases.append(model_to_use.forwarded_model_id)
for alias in aliases:
_add_candidate(alias, model_to_use, upstream_for_override)
seen_model_provider.add(dedupe_key)

View File

@@ -7,11 +7,12 @@ from datetime import datetime
from typing import Optional
from fastapi import HTTPException
from sqlalchemy import case
from sqlalchemy.exc import IntegrityError
from sqlmodel import col, select, update
from .core import get_logger
from .core.db import ApiKey, AsyncSession
from .core.db import ApiKey, AsyncSession, accumulate_routstr_fee
from .core.settings import settings
from .payment.cost_calculation import (
CostData,
@@ -22,6 +23,13 @@ from .payment.cost_calculation import (
from .wallet import credit_balance, deserialize_token_from_string
logger = get_logger(__name__)
payments_logger = get_logger("routstr.payments")
# Routstr platform fee constants
ROUTSTR_FEE_PERCENT: float = 2.1
ROUTSTR_LN_ADDRESS: str = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900
ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200
# TODO: implement prepaid api key (not like it was before)
# PREPAID_API_KEY = os.environ.get("PREPAID_API_KEY", None)
@@ -584,6 +592,18 @@ async def pay_for_request(
"total_requests": billing_key.total_requests,
},
)
payments_logger.info(
"RESERVE",
extra={
"event": "reserve",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost_reserved": cost_per_request,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
},
)
return cost_per_request
@@ -635,6 +655,17 @@ async def revert_pay_for_request(
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
payments_logger.info(
"REVERT",
extra={
"event": "revert",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost_reverted": cost_per_request,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
},
)
return True
@@ -716,6 +747,17 @@ async def adjust_payment_for_tokens(
},
)
async def _accumulate_fee(total_cost_msats: int) -> None:
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",
extra={"error": str(e), "fee_msats": fee_msats},
)
match await calculate_cost(response_data, deducted_max_cost, session):
case MaxCostData() as cost:
logger.debug(
@@ -728,11 +770,32 @@ async def adjust_payment_for_tokens(
},
)
# Finalize by releasing reservation and charging max cost
if billing_key.reserved_balance < deducted_max_cost:
logger.error(
"reserved_balance below deducted_max_cost before MaxCost finalization — clamping to 0",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"reserved_balance": billing_key.reserved_balance,
"deducted_max_cost": deducted_max_cost,
"total_cost_msats": cost.total_msats,
"balance": billing_key.balance,
"total_spent": billing_key.total_spent,
"model": model,
},
)
safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost,
reserved_balance=safe_reserved,
balance=col(ApiKey.balance) - cost.total_msats,
total_spent=col(ApiKey.total_spent) + cost.total_msats,
)
@@ -741,13 +804,17 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
total_spent=col(ApiKey.total_spent) + cost.total_msats,
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=child_safe_reserved,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -782,6 +849,24 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(cost.total_msats)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": cost.total_msats,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "max_cost",
},
)
return cost.dict()
case CostData() as cost:
@@ -815,12 +900,32 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
if billing_key.reserved_balance < deducted_max_cost:
logger.error(
"reserved_balance below deducted_max_cost on exact-cost finalization — clamping to 0",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"reserved_balance": billing_key.reserved_balance,
"deducted_max_cost": deducted_max_cost,
"total_cost_msats": total_cost_msats,
"balance": billing_key.balance,
"total_spent": billing_key.total_spent,
"model": model,
},
)
exact_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=exact_safe_reserved,
balance=col(ApiKey.balance) - total_cost_msats,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
)
@@ -829,13 +934,17 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_exact_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
total_spent=col(ApiKey.total_spent) + total_cost_msats,
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=child_exact_safe_reserved,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -844,44 +953,56 @@ async def adjust_payment_for_tokens(
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
return cost.dict()
# this should never happen why do we handle this???
if cost_difference > 0:
# Need to charge more than reserved, finalize by releasing reservation and charging total
logger.info(
"Additional charge required for token usage",
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"additional_charge": cost_difference,
"current_balance": billing_key.balance,
"sufficient_balance": billing_key.balance >= cost_difference,
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": total_cost_msats,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "exact",
},
)
return cost.dict()
# actual cost exceeded discounted reservation (due to tolerance_percentage)
if cost_difference > 0:
# Always release the reservation and charge min(actual_cost, balance).
# Using a CASE expression makes this a single atomic UPDATE — no
# multi-level fallback needed and balance can never go negative.
chargeable = case(
(col(ApiKey.balance) >= total_cost_msats, total_cost_msats),
else_=col(ApiKey.balance),
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
balance=col(ApiKey.balance) - total_cost_msats,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost,
balance=col(ApiKey.balance) - chargeable,
total_spent=col(ApiKey.total_spent) + chargeable,
)
)
result = await session.exec(finalize_stmt) # type: ignore[call-overload]
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
.values(
total_spent=col(ApiKey.total_spent) + total_cost_msats,
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost,
total_spent=col(ApiKey.total_spent) + min(billing_key.balance, total_cost_msats),
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -889,11 +1010,10 @@ async def adjust_payment_for_tokens(
await session.commit()
if result.rowcount:
cost.total_msats = total_cost_msats
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
cost.total_msats = total_cost_msats
logger.info(
"Finalized payment with additional charge",
extra={
@@ -904,9 +1024,29 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": total_cost_msats,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "overrun",
},
)
else:
# Guard fired: reservation was already released by a concurrent
# finalization for this key. Nothing left to do.
logger.warning(
"Failed to finalize additional charge - releasing reservation",
"Finalization skipped - reservation already released",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
@@ -914,7 +1054,6 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await release_reservation_only()
else:
# Refund some of the base cost
refund = abs(cost_difference)
@@ -929,12 +1068,33 @@ async def adjust_payment_for_tokens(
},
)
if billing_key.reserved_balance < deducted_max_cost:
logger.error(
"reserved_balance below deducted_max_cost on refund finalization — clamping to 0",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"reserved_balance": billing_key.reserved_balance,
"deducted_max_cost": deducted_max_cost,
"total_cost_msats": total_cost_msats,
"refund_amount": refund,
"balance": billing_key.balance,
"total_spent": billing_key.total_spent,
"model": model,
},
)
refund_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
refund_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=refund_safe_reserved,
balance=col(ApiKey.balance) - total_cost_msats,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
)
@@ -943,13 +1103,17 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_refund_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
total_spent=col(ApiKey.total_spent) + total_cost_msats,
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=child_refund_safe_reserved,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -986,6 +1150,25 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": total_cost_msats,
"refunded": refund,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "refund",
},
)
return cost.dict()

View File

@@ -7,10 +7,16 @@ from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel
from sqlmodel import select
from sqlmodel import col, select, update
from .auth import get_billing_key, validate_bearer_key
from .core.db import ApiKey, AsyncSession, CashuTransaction, get_session
from .core.db import (
ApiKey,
AsyncSession,
CashuTransaction,
get_session,
store_cashu_transaction,
)
from .core.logging import get_logger
from .core.settings import settings
from .lightning import lightning_router
@@ -205,6 +211,26 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None:
_refund_cache[key] = (expiry, value)
async def _restore_balance(
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int
) -> 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},
)
@router.post("/refund", response_model=None)
async def refund_wallet_endpoint(
authorization: Annotated[str | None, Header()] = None,
@@ -286,7 +312,31 @@ async def refund_wallet_endpoint(
elif remaining_balance <= 0:
raise HTTPException(status_code=400, detail="No balance to refund")
# Perform refund operation first, before modifying balance
# Capture values before debit — the session may refresh key after commit
pre_debit_balance = key.balance
pre_debit_reserved = key.reserved_balance
# --- DEBIT FIRST: atomically zero the balance before minting tokens ---
# This prevents the race where a concurrent topup/spend happens between
# reading the balance and minting the refund token (double-spend).
debit_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.balance) == pre_debit_balance)
.where(col(ApiKey.reserved_balance) == pre_debit_reserved)
.values(balance=0, reserved_balance=0)
)
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.",
)
# --- MINT: balance is locked at zero, safe to create the refund token ---
try:
if key.refund_address:
from .core.settings import settings as global_settings
@@ -310,11 +360,24 @@ async def refund_wallet_endpoint(
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 HTTPException:
# Re-raise HTTP exceptions (like 400 for balance too small)
# Minting failed — restore the debited balance
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved)
raise
except Exception as e:
# If refund fails, don't modify the database
# Minting failed — restore the debited balance
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved)
error_msg = str(e)
if (
"mint" in error_msg.lower()
@@ -328,23 +391,67 @@ async def refund_wallet_endpoint(
await _refund_cache_set(bearer_value, result)
previous_reserved_balance = key.reserved_balance
key.balance = 0
key.reserved_balance = 0
session.add(key)
await session.commit()
if "token" in result:
try:
await store_cashu_transaction(
token=result["token"],
amount=remaining_balance,
unit=key.refund_currency or "sat",
mint_url=key.refund_mint_url,
typ="out",
collected=False,
source="apikey",
api_key_hashed_key=key.hashed_key,
)
except Exception:
pass # store_cashu_transaction already logs
logger.info(
"refund_wallet_endpoint: refund successful",
extra={
"refunded_msats": remaining_balance_msats,
"previous_reserved_balance": previous_reserved_balance,
"previous_reserved_balance": key.reserved_balance,
},
)
return result
@router.get("/history")
async def wallet_history(
key: ApiKey = Depends(get_key_from_header),
session: AsyncSession = Depends(get_session),
) -> dict[str, list[dict[str, str | int | bool | None]]]:
if key.parent_key_hash:
raise HTTPException(
status_code=400,
detail="Cannot view child key history. Please use the parent key instead.",
)
result = await session.exec(
select(CashuTransaction)
.where(CashuTransaction.api_key_hashed_key == key.hashed_key)
.order_by(col(CashuTransaction.created_at).desc())
)
transactions = result.all()
return {
"transactions": [
{
"id": tx.id,
"type": tx.type,
"source": tx.source,
"amount": tx.amount,
"unit": tx.unit,
"mint_url": tx.mint_url,
"created_at": tx.created_at,
"collected": tx.collected,
"swept": tx.swept,
}
for tx in transactions
]
}
@router.post("/donate")
async def donate(token: str, ref: str | None = None) -> str:
try:

View File

@@ -9,7 +9,7 @@ from pydantic import BaseModel
from sqlmodel import select
from ..payment.models import _row_to_model, list_models
from ..proxy import refresh_model_maps, reinitialize_upstreams
from ..proxy import refresh_model_maps, reinitialize_upstreams, sync_provider_fees
from ..wallet import (
fetch_all_balances,
get_proofs_per_mint_and_unit,
@@ -20,6 +20,7 @@ from ..wallet import (
from .db import (
ApiKey,
CashuTransaction,
CliToken,
ModelRow,
UpstreamProviderRow,
create_session,
@@ -38,12 +39,27 @@ ADMIN_SESSION_DURATION = 3600
MAX_USAGE_ANALYTICS_HOURS = 365 * 24
def require_admin_api(request: Request) -> None:
async def require_admin_api(request: Request) -> None:
auth_header = request.headers.get("Authorization")
if auth_header and auth_header.startswith("Bearer "):
token = auth_header.split(" ", 1)[1]
expiry = admin_sessions.get(token)
if expiry and expiry > int(datetime.now(timezone.utc).timestamp()):
if not auth_header or not auth_header.startswith("Bearer "):
raise HTTPException(status_code=403, detail="Unauthorized")
token = auth_header.split(" ", 1)[1]
now_ts = int(datetime.now(timezone.utc).timestamp())
# 1) Short-lived session token (in-memory)
expiry = admin_sessions.get(token)
if expiry and expiry > now_ts:
return
# 2) Long-lived CLI token (DB-backed)
async with create_session() as session:
result = await session.exec(select(CliToken).where(CliToken.token == token))
cli_token = result.first()
if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts):
cli_token.last_used_at = now_ts
session.add(cli_token)
await session.commit()
return
raise HTTPException(status_code=403, detail="Unauthorized")
@@ -242,6 +258,73 @@ async def admin_logout(request: Request) -> dict[str, object]:
return {"ok": True}
# ─── CLI Tokens (long-lived bearer tokens for CLI/agent use) ───
class CliTokenCreate(BaseModel):
name: str
expires_in_days: int | None = None
@admin_router.get("/api/cli-tokens", dependencies=[Depends(require_admin_api)])
async def list_cli_tokens() -> list[dict[str, object]]:
async with create_session() as session:
result = await session.exec(select(CliToken))
tokens = result.all()
return [
{
"id": t.id,
"name": t.name,
"token_preview": f"{t.token[:8]}...{t.token[-4:]}",
"created_at": t.created_at,
"last_used_at": t.last_used_at,
"expires_at": t.expires_at,
}
for t in tokens
]
@admin_router.post("/api/cli-tokens", dependencies=[Depends(require_admin_api)])
async def create_cli_token(payload: CliTokenCreate) -> dict[str, object]:
name = (payload.name or "").strip()
if not name:
raise HTTPException(status_code=400, detail="Name is required")
raw_token = secrets.token_urlsafe(32)
expires_at: int | None = None
if payload.expires_in_days is not None and payload.expires_in_days > 0:
expires_at = int(datetime.now(timezone.utc).timestamp()) + (
payload.expires_in_days * 86400
)
async with create_session() as session:
cli_token = CliToken(token=raw_token, name=name, expires_at=expires_at)
session.add(cli_token)
await session.commit()
await session.refresh(cli_token)
return {
"id": cli_token.id,
"name": cli_token.name,
"token": raw_token, # full token returned only on creation
"created_at": cli_token.created_at,
"expires_at": cli_token.expires_at,
}
@admin_router.delete(
"/api/cli-tokens/{token_id}", dependencies=[Depends(require_admin_api)]
)
async def revoke_cli_token(token_id: str) -> dict[str, object]:
async with create_session() as session:
cli_token = await session.get(CliToken, token_id)
if not cli_token:
raise HTTPException(status_code=404, detail="Token not found")
await session.delete(cli_token)
await session.commit()
return {"ok": True, "deleted_id": token_id}
class WithdrawRequest(BaseModel):
amount: int
mint_url: str | None = None
@@ -295,6 +378,7 @@ class ModelCreate(BaseModel):
canonical_slug: str | None = None
alias_ids: list[str] | None = None
enabled: bool = True
forwarded_model_id: str | None = None
@admin_router.post(
@@ -339,6 +423,7 @@ async def upsert_provider_model(
json.dumps(payload.alias_ids) if payload.alias_ids else None
)
existing_row.enabled = payload.enabled
existing_row.forwarded_model_id = payload.forwarded_model_id or payload.id
session.add(existing_row)
await session.commit()
@@ -371,6 +456,7 @@ async def upsert_provider_model(
),
upstream_provider_id=provider_id,
enabled=payload.enabled,
forwarded_model_id=payload.forwarded_model_id or payload.id,
)
session.add(row)
await session.commit()
@@ -553,6 +639,7 @@ class UpstreamProviderCreate(BaseModel):
api_version: str | None = None
enabled: bool = True
provider_fee: float = 1.01
provider_fee_default: float | None = None
provider_settings: dict | None = None
@@ -563,29 +650,37 @@ class UpstreamProviderUpdate(BaseModel):
api_version: str | None = None
enabled: bool | None = None
provider_fee: float | None = None
provider_fee_default: float | None = None
provider_settings: dict | None = None
def _provider_to_dict(
p: UpstreamProviderRow, redact_key: bool = True
) -> dict[str, object]:
return {
"id": p.id,
"provider_type": p.provider_type,
"base_url": p.base_url,
"api_key": "[REDACTED]" if (redact_key and p.api_key) else (p.api_key or ""),
"api_version": p.api_version,
"enabled": p.enabled,
"provider_fee": p.provider_fee,
"provider_fee_default": p.provider_fee_default,
"provider_settings": json.loads(p.provider_settings)
if p.provider_settings
else None,
"provider_fee_schedules": json.loads(p.provider_fee_schedules)
if p.provider_fee_schedules
else [],
}
@admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)])
async def get_upstream_providers() -> list[dict[str, object]]:
async with create_session() as session:
result = await session.exec(select(UpstreamProviderRow))
providers = result.all()
return [
{
"id": p.id,
"provider_type": p.provider_type,
"base_url": p.base_url,
"api_key": "[REDACTED]" if p.api_key else "",
"api_version": p.api_version,
"enabled": p.enabled,
"provider_fee": p.provider_fee,
"provider_settings": json.loads(p.provider_settings)
if p.provider_settings
else None,
}
for p in providers
]
return [_provider_to_dict(p) for p in providers]
@admin_router.post("/api/upstream-providers", dependencies=[Depends(require_admin_api)])
@@ -612,6 +707,9 @@ async def create_upstream_provider(
api_version=payload.api_version,
enabled=payload.enabled,
provider_fee=payload.provider_fee,
provider_fee_default=payload.provider_fee_default
if payload.provider_fee_default is not None
else payload.provider_fee,
provider_settings=json.dumps(payload.provider_settings)
if payload.provider_settings
else None,
@@ -621,17 +719,7 @@ async def create_upstream_provider(
await session.refresh(provider)
await reinitialize_upstreams()
await refresh_model_maps()
return {
"id": provider.id,
"provider_type": provider.provider_type,
"base_url": provider.base_url,
"api_key": "[REDACTED]",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
"provider_settings": payload.provider_settings,
}
return _provider_to_dict(provider)
@admin_router.get(
@@ -642,18 +730,7 @@ async def get_upstream_provider(provider_id: int) -> dict[str, object]:
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
return {
"id": provider.id,
"provider_type": provider.provider_type,
"base_url": provider.base_url,
"api_key": "[REDACTED]" if provider.api_key else "",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
"provider_settings": json.loads(provider.provider_settings)
if provider.provider_settings
else None,
}
return _provider_to_dict(provider)
@admin_router.patch(
@@ -679,6 +756,8 @@ async def update_upstream_provider(
provider.enabled = payload.enabled
if payload.provider_fee is not None:
provider.provider_fee = payload.provider_fee
if payload.provider_fee_default is not None:
provider.provider_fee_default = payload.provider_fee_default
if payload.provider_settings is not None:
provider.provider_settings = json.dumps(payload.provider_settings)
@@ -687,19 +766,7 @@ async def update_upstream_provider(
await session.refresh(provider)
await reinitialize_upstreams()
await refresh_model_maps()
return {
"id": provider.id,
"provider_type": provider.provider_type,
"base_url": provider.base_url,
"api_key": "[REDACTED]",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
"provider_settings": json.loads(provider.provider_settings)
if provider.provider_settings
else None,
}
return _provider_to_dict(provider)
@admin_router.delete(
@@ -717,6 +784,78 @@ async def delete_upstream_provider(provider_id: int) -> dict[str, object]:
return {"ok": True, "deleted_id": provider_id}
@admin_router.get(
"/api/upstream-providers/{provider_id}/fee-schedules",
dependencies=[Depends(require_admin_api)],
)
async def get_fee_schedules(provider_id: int) -> list[dict]:
async with create_session() as session:
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
return (
json.loads(provider.provider_fee_schedules)
if provider.provider_fee_schedules
else []
)
class FeeScheduleUpdate(BaseModel):
schedules: list[dict]
@admin_router.put(
"/api/upstream-providers/{provider_id}/fee-schedules",
dependencies=[Depends(require_admin_api)],
)
async def update_fee_schedules(
provider_id: int, payload: FeeScheduleUpdate
) -> list[dict]:
from ..payment.fee_schedule import FeeTimeRange, validate_no_overlaps
try:
ranges = [FeeTimeRange(**s) for s in payload.schedules]
except Exception as e:
raise HTTPException(status_code=400, detail=f"Invalid schedule data: {e}")
try:
validate_no_overlaps(ranges)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
serialized = [r.dict() for r in ranges]
async with create_session() as session:
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
provider.provider_fee_schedules = json.dumps(serialized)
session.add(provider)
await session.commit()
await sync_provider_fees()
await refresh_model_maps()
return serialized
@admin_router.delete(
"/api/upstream-providers/{provider_id}/fee-schedules",
dependencies=[Depends(require_admin_api)],
)
async def delete_fee_schedules(provider_id: int) -> dict:
async with create_session() as session:
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
provider.provider_fee_schedules = None
session.add(provider)
await session.commit()
await sync_provider_fees()
await refresh_model_maps()
return {"ok": True}
@admin_router.get("/api/provider-types", dependencies=[Depends(require_admin_api)])
async def get_provider_types() -> list[dict[str, object]]:
"""Get metadata about available provider types including default URLs and whether they're fixed."""
@@ -1332,41 +1471,56 @@ async def get_transactions_api(
type: str | None = None,
status: str | None = None,
search: str | None = None,
limit: int = 100,
source: str | None = None,
limit: int = 50,
offset: int = 0,
) -> dict:
async with create_session() as session:
from sqlmodel import col
from sqlmodel import col, func
stmt = select(CashuTransaction)
base = select(CashuTransaction)
if type:
stmt = stmt.where(CashuTransaction.type == type)
base = base.where(CashuTransaction.type == type)
if source:
if source == "x-cashu":
base = base.where(
(CashuTransaction.source == "x-cashu")
| (CashuTransaction.source == None) # noqa: E711
)
else:
base = base.where(CashuTransaction.source == source)
if status:
if status == "collected":
stmt = stmt.where(CashuTransaction.collected == True) # noqa: E712
base = base.where(CashuTransaction.collected == True) # noqa: E712
elif status == "swept":
stmt = stmt.where(CashuTransaction.swept == True) # noqa: E712
base = base.where(CashuTransaction.swept == True) # noqa: E712
elif status == "pending":
stmt = stmt.where(
base = base.where(
CashuTransaction.collected == False, # noqa: E712
CashuTransaction.swept == False, # noqa: E712
)
if search:
search_pattern = f"%{search}%"
stmt = stmt.where(
base = base.where(
(col(CashuTransaction.id).like(search_pattern))
| (col(CashuTransaction.token).like(search_pattern))
| (col(CashuTransaction.request_id).like(search_pattern))
| (col(CashuTransaction.api_key_hashed_key).like(search_pattern))
)
stmt = stmt.order_by(col(CashuTransaction.created_at).desc()).limit(limit)
count_result = await session.exec(
select(func.count()).select_from(base.subquery())
)
total = count_result.one()
stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit)
results = await session.exec(stmt)
transactions = results.all()
return {
"transactions": [tx.dict() for tx in transactions],
"total": len(transactions),
"total": total,
}

View File

@@ -10,8 +10,9 @@ from alembic import command
from alembic.config import Config
from alembic.util.exc import CommandError
from sqlalchemy import UniqueConstraint
from sqlalchemy.exc import OperationalError
from sqlalchemy.ext.asyncio.engine import create_async_engine
from sqlmodel import Field, Relationship, SQLModel, func, select, update
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
from sqlmodel.ext.asyncio.session import AsyncSession
from .logging import get_logger
@@ -105,6 +106,10 @@ class ModelRow(SQLModel, table=True): # type: ignore
default=None, description="JSON array of model alias IDs"
)
enabled: bool = Field(default=True, description="Whether this model is enabled")
forwarded_model_id: str | None = Field(
default=None,
description="Model ID to use when forwarding requests to upstream provider. Defaults to id if not set.",
)
upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models")
@@ -150,6 +155,16 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
)
collected: bool = Field(default=False)
swept: bool = Field(default=False)
source: str = Field(
default="x-cashu",
description="Payment source: x-cashu or apikey",
)
api_key_hashed_key: str | None = Field(
default=None,
foreign_key="api_keys.hashed_key",
index=True,
description="Associated API key hash for wallet history",
)
async def store_cashu_transaction(
@@ -161,6 +176,8 @@ async def store_cashu_transaction(
request_id: str | None = None,
collected: bool = False,
created_at: int | None = None,
source: str = "x-cashu",
api_key_hashed_key: str | None = None,
) -> None:
try:
async with create_session() as session:
@@ -173,6 +190,8 @@ async def store_cashu_transaction(
request_id=request_id,
collected=collected,
created_at=created_at or int(time.time()),
source=source,
api_key_hashed_key=api_key_hashed_key,
)
session.add(tx)
await session.commit()
@@ -201,17 +220,83 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
)
enabled: bool = Field(default=True, description="Whether this provider is enabled")
provider_fee: float = Field(
default=1.01, description="Provider fee multiplier (default 1%)"
default=1.01, description="Active fee multiplier (can be set by schedule)"
)
provider_fee_default: float = Field(
default=1.01, description="Default fee multiplier (outside schedules)"
)
provider_settings: str | None = Field(
default=None, description="JSON string for provider-specific settings"
)
provider_fee_schedules: str | None = Field(
default=None, description="JSON array of fee time ranges (HH:MM UTC)"
)
models: list["ModelRow"] = Relationship(
back_populates="upstream_provider",
sa_relationship_kwargs={"cascade": "all, delete-orphan"},
)
class RoutstrFee(SQLModel, table=True): # type: ignore
__tablename__ = "routstr_fees"
id: int = Field(default=1, primary_key=True)
accumulated_msats: int = Field(default=0)
total_paid_msats: int = Field(default=0)
last_paid_at: int | None = Field(default=None)
class CliToken(SQLModel, table=True): # type: ignore
"""Long-lived authorization token for CLI/agent use against admin endpoints."""
__tablename__ = "cli_tokens"
id: str = Field(
primary_key=True, default_factory=lambda: uuid.uuid4().hex
)
token: str = Field(unique=True, index=True, description="Bearer token value")
name: str = Field(description="Human-readable label for this token")
created_at: int = Field(default_factory=lambda: int(time.time()))
last_used_at: int | None = Field(default=None)
expires_at: int | None = Field(
default=None, description="Optional expiry unix timestamp; null = never expires"
)
async def accumulate_routstr_fee(session: AsyncSession, amount_msats: int) -> None:
stmt = (
update(RoutstrFee)
.where(col(RoutstrFee.id) == 1)
.values(accumulated_msats=RoutstrFee.accumulated_msats + amount_msats)
)
result = await session.exec(stmt) # type: ignore[call-overload]
if result.rowcount == 0:
session.add(RoutstrFee(id=1, accumulated_msats=amount_msats))
await session.commit()
async def get_routstr_fee(session: AsyncSession) -> RoutstrFee:
fee = await session.get(RoutstrFee, 1)
if fee is None:
fee = RoutstrFee(id=1, accumulated_msats=0, total_paid_msats=0)
session.add(fee)
await session.commit()
await session.refresh(fee)
return fee
async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> None:
stmt = (
update(RoutstrFee)
.where(col(RoutstrFee.id) == 1)
.values(
accumulated_msats=RoutstrFee.accumulated_msats - paid_msats,
total_paid_msats=RoutstrFee.total_paid_msats + paid_msats,
last_paid_at=int(time.time()),
)
)
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
async def balances_for_mint_and_unit(
db_session: AsyncSession, mint_url: str, unit: str
) -> int:
@@ -327,6 +412,17 @@ def run_migrations() -> None:
command.stamp(alembic_cfg, "head")
else:
raise
except OperationalError as e:
if "duplicate column name" in str(e).lower():
logger.warning(
"Migration hit a column that already exists (likely added via "
"create_all on another branch). Stamping to current head.",
extra={"error": str(e)},
)
_clear_alembic_version()
command.stamp(alembic_cfg, "head")
else:
raise
logger.info("Database migrations completed successfully")

View File

@@ -22,7 +22,7 @@ from ..payment.models import models_router, update_sats_pricing
from ..payment.price import update_prices_periodically
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
from ..upstream.auto_topup import periodic_auto_topup
from ..wallet import periodic_payout, periodic_refund_sweep
from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout
from .admin import admin_router
from .db import create_session, init_db, run_migrations
from .exceptions import general_exception_handler, http_exception_handler
@@ -36,9 +36,9 @@ setup_logging()
logger = get_logger(__name__)
if os.getenv("VERSION_SUFFIX") is not None:
__version__ = f"0.4.1-{os.getenv('VERSION_SUFFIX')}"
__version__ = f"0.4.3-{os.getenv('VERSION_SUFFIX')}"
else:
__version__ = "0.4.1"
__version__ = "0.4.3"
@asynccontextmanager
@@ -56,6 +56,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
key_reset_task = None
auto_topup_task = None
refund_sweep_task = None
routstr_fee_task = None
try:
# Run database migrations on startup
@@ -102,8 +103,11 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
btc_price_task = asyncio.create_task(update_prices_periodically())
pricing_task = asyncio.create_task(update_sats_pricing())
if global_settings.models_refresh_interval_seconds > 0:
# Pass the accessor (not its current value) so the loop sees providers
# added/changed via reinitialize_upstreams() instead of staying pinned
# to the startup snapshot.
models_refresh_task = asyncio.create_task(
refresh_upstreams_models_periodically(get_upstreams())
refresh_upstreams_models_periodically(get_upstreams)
)
model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically())
payout_task = asyncio.create_task(periodic_payout())
@@ -115,6 +119,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
key_reset_task = asyncio.create_task(periodic_key_reset())
auto_topup_task = asyncio.create_task(periodic_auto_topup())
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout())
yield
@@ -152,6 +157,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
auto_topup_task.cancel()
if refund_sweep_task is not None:
refund_sweep_task.cancel()
if routstr_fee_task is not None:
routstr_fee_task.cancel()
try:
tasks_to_wait = []
@@ -177,6 +184,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(auto_topup_task)
if refund_sweep_task is not None:
tasks_to_wait.append(refund_sweep_task)
if routstr_fee_task is not None:
tasks_to_wait.append(routstr_fee_task)
if tasks_to_wait:
await asyncio.gather(*tasks_to_wait, return_exceptions=True)

View File

@@ -38,11 +38,6 @@ class LoggingMiddleware(BaseHTTPMiddleware):
except Exception:
pass
# Extract request info
client_host = None
if request.client:
client_host = request.client.host
# Log incoming request
logger.info(
"Incoming request",
@@ -51,7 +46,6 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"method": request.method,
"path": request.url.path,
"query_params": dict(request.query_params),
"client_host": client_host,
"headers": {
k: v
for k, v in request.headers.items()
@@ -100,7 +94,6 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"path": request.url.path,
"status_code": response.status_code,
"duration_ms": round(duration * 1000, 2),
"client_host": client_host,
},
)
if hasattr(response, "headers"):
@@ -120,7 +113,6 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"method": request.method,
"path": request.url.path,
"duration_ms": round(duration * 1000, 2),
"client_host": client_host,
"error": str(e),
"error_type": type(e).__name__,
},

View File

@@ -74,7 +74,7 @@ class Settings(BaseSettings):
enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH")
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(default=86400, env="REFUND_SWEEP_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS")
# Logging
log_level: str = Field(default="INFO", env="LOG_LEVEL")

View File

@@ -0,0 +1,113 @@
"""Dynamic provider fee schedule logic.
Supports time-based fee ranges (HH:MM UTC) with overlap validation and active fee resolution.
"""
from __future__ import annotations
import re
from datetime import datetime, timezone
from pydantic.v1 import BaseModel, validator
_HH_MM_RE = re.compile(r"^([01]\d|2[0-3]):([0-5]\d)$")
class FeeTimeRange(BaseModel):
start_time: str # HH:MM UTC
end_time: str # HH:MM UTC
provider_fee: float
@validator("start_time", "end_time")
@classmethod
def validate_time_format(cls, v: str) -> str:
if not _HH_MM_RE.match(v):
raise ValueError(f"Time must be in HH:MM format (00:0023:59), got: {v!r}")
return v
@validator("provider_fee")
@classmethod
def validate_fee(cls, v: float) -> float:
if v <= 0:
raise ValueError(f"provider_fee must be > 0 (got {v})")
return v
def _to_minutes(t: str) -> int:
h, m = map(int, t.split(":"))
return h * 60 + m
def _range_intervals(r: FeeTimeRange) -> list[tuple[int, int]]:
"""Return list of [start, end) minute intervals for this range.
Handles midnight-crossing (e.g. 22:0006:00 → [(1320,1440),(0,360)]).
start == end is treated as a full-day range.
"""
start = _to_minutes(r.start_time)
end = _to_minutes(r.end_time)
if start < end:
return [(start, end)]
if start > end:
return [(start, 1440), (0, end)]
# start == end → full day
return [(0, 1440)]
def _intervals_overlap(a: tuple[int, int], b: tuple[int, int]) -> bool:
return a[0] < b[1] and b[0] < a[1]
def ranges_overlap(a: FeeTimeRange, b: FeeTimeRange) -> bool:
"""Return True if two fee time ranges overlap at any point in the day."""
for ia in _range_intervals(a):
for ib in _range_intervals(b):
if _intervals_overlap(ia, ib):
return True
return False
def validate_no_overlaps(ranges: list[FeeTimeRange]) -> None:
"""Raise ValueError if any two ranges in the list overlap."""
for i in range(len(ranges)):
for j in range(i + 1, len(ranges)):
if ranges_overlap(ranges[i], ranges[j]):
raise ValueError(
f"Fee ranges overlap: [{ranges[i].start_time}{ranges[i].end_time}]"
f" and [{ranges[j].start_time}{ranges[j].end_time}]"
)
def get_active_fee(
ranges: list[FeeTimeRange] | None,
default_fee: float,
*,
_now: datetime | None = None,
) -> float:
"""Return the provider fee for the current UTC time.
Falls back to *default_fee* when no range matches or *ranges* is empty/None.
The *_now* parameter is for testing only.
"""
if not ranges or not isinstance(ranges, list):
return default_fee
now = _now if _now is not None else datetime.now(timezone.utc)
# Normalize to UTC
if now.tzinfo is not None:
now = now.astimezone(timezone.utc)
current = now.hour * 60 + now.minute
for r in ranges:
start = _to_minutes(r.start_time)
end = _to_minutes(r.end_time)
if start < end:
if start <= current < end:
return r.provider_fee
elif start > end: # midnight-crossing
if current >= start or current < end:
return r.provider_fee
else: # full day (start == end)
return r.provider_fee
return default_fee

View File

@@ -7,7 +7,7 @@ from fastapi import APIRouter, Depends
from pydantic.v1 import BaseModel
from sqlmodel.ext.asyncio.session import AsyncSession
from ..core.db import ModelRow, get_session
from ..core.db import ModelRow, UpstreamProviderRow, get_session
from ..core.logging import get_logger
from ..core.settings import settings
from .price import sats_usd_price
@@ -60,6 +60,7 @@ class Model(BaseModel):
upstream_provider_id: int | str | None = None
canonical_slug: str | None = None
alias_ids: list[str] | None = None
forwarded_model_id: str | None = None
def __hash__(self) -> int:
return hash(self.id)
@@ -177,6 +178,7 @@ def _row_to_model(
upstream_provider_id=row.upstream_provider_id,
canonical_slug=getattr(row, "canonical_slug", None),
alias_ids=json.loads(row.alias_ids) if row.alias_ids else None,
forwarded_model_id=getattr(row, "forwarded_model_id", None) or row.id,
)
if apply_provider_fee:
@@ -329,6 +331,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
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:
logger.error(
@@ -402,6 +405,76 @@ async def update_sats_pricing() -> None:
logger.error(f"Error updating sats pricing: {e}")
class ModelTestRequest(BaseModel):
model_id: str
endpoint_type: str
request_data: dict
@models_router.post("/api/models/test")
async def test_model(
payload: ModelTestRequest,
session: AsyncSession = Depends(get_session),
) -> dict:
"""Test a model by sending a request through its configured upstream provider."""
from sqlmodel import select
result = await session.execute(
select(ModelRow).where(ModelRow.id == payload.model_id)
)
model_row = result.scalars().first()
if not model_row:
return {
"success": False,
"error": f"Model '{payload.model_id}' not found in database",
"status_code": 404,
}
provider = await session.get(UpstreamProviderRow, model_row.upstream_provider_id)
if not provider:
return {
"success": False,
"error": "Upstream provider not found",
"status_code": 404,
}
base_url = provider.base_url.rstrip("/")
if payload.endpoint_type == "chat-completions":
url = f"{base_url}/chat/completions"
else:
url = f"{base_url}/{payload.endpoint_type}"
actual_model_id = model_row.forwarded_model_id or model_row.id
request_data = dict(payload.request_data)
request_data["model"] = actual_model_id
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {provider.api_key}",
}
try:
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.post(url, json=request_data, headers=headers)
try:
response_data = response.json()
except Exception:
response_data = {"raw": response.text}
return {
"success": response.status_code < 400,
"data": response_data,
"status_code": response.status_code,
}
except Exception as e:
return {
"success": False,
"error": str(e),
"status_code": 500,
}
@models_router.get("/v1/models")
@models_router.get("/v1/models/", include_in_schema=False)
@models_router.get("/models")
@@ -411,4 +484,10 @@ async def models(session: AsyncSession = Depends(get_session)) -> dict:
from ..proxy import get_unique_models
items = get_unique_models()
return {"data": items}
data = []
for model in items:
m = model.dict()
if model.forwarded_model_id:
m["id"] = model.forwarded_model_id
data.append(m)
return {"data": data}

View File

@@ -44,6 +44,7 @@ async def initialize_upstreams() -> None:
global _upstreams
_upstreams = await init_upstreams()
logger.info(f"Initialized {len(_upstreams)} upstream providers")
await sync_provider_fees()
await refresh_model_maps()
@@ -55,6 +56,7 @@ async def reinitialize_upstreams() -> None:
"Re-initialized upstream providers from admin action",
extra={"provider_count": len(_upstreams)},
)
await sync_provider_fees()
await refresh_model_maps()
@@ -69,7 +71,25 @@ def get_upstreams() -> list[BaseUpstreamProvider]:
def get_model_instance(model_id: str) -> Model | None:
"""Get Model instance by ID from global cache."""
return _model_instances.get(model_id.lower())
if not model_id:
return None
model_id_lower = model_id.lower()
# Try exact match first
if model := _model_instances.get(model_id_lower):
return model
# Try stripping common version suffixes (e.g., -20251222)
# This handles cases where upstream returns a specific version
# but we only track the base model name.
import re
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
if base_model_id != model_id_lower:
if model := _model_instances.get(base_model_id):
return model
return None
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
@@ -100,6 +120,12 @@ async def refresh_model_maps() -> None:
disabled_model_ids: set[str] = set()
for provider in provider_rows:
# Match with instance in _upstreams to update its state from DB
for upstream in _upstreams:
if getattr(upstream, "db_id", None) == provider.id:
# This updates fee and merges DB models WITHOUT hitting network
await upstream.refresh_models_cache(skip_network=True)
if not provider.enabled:
continue
for model in provider.models:
@@ -115,6 +141,39 @@ async def refresh_model_maps() -> None:
)
async def sync_provider_fees() -> None:
"""Update active provider_fee in database based on schedules and defaults."""
from .payment.fee_schedule import FeeTimeRange, get_active_fee
async with create_session() as session:
result = await session.exec(select(UpstreamProviderRow))
provider_rows = result.all()
updated = False
for p in provider_rows:
schedules = None
if p.provider_fee_schedules:
try:
schedules = [
FeeTimeRange(**s) for s in json.loads(p.provider_fee_schedules)
]
except Exception:
pass
active_fee = get_active_fee(schedules, p.provider_fee_default)
if p.provider_fee != active_fee:
logger.info(
f"Updating active fee for provider {p.id}: {p.provider_fee} -> {active_fee}",
extra={"provider_id": p.id, "active_fee": active_fee},
)
p.provider_fee = active_fee
session.add(p)
updated = True
if updated:
await session.commit()
async def refresh_model_maps_periodically() -> None:
"""Background task to refresh model maps every minute."""
import asyncio
@@ -122,6 +181,7 @@ async def refresh_model_maps_periodically() -> None:
while True:
try:
await asyncio.sleep(60)
await sync_provider_fees()
await refresh_model_maps()
except asyncio.CancelledError:
break
@@ -207,7 +267,7 @@ async def proxy(
elif auth := headers.get("authorization", None):
key = await get_bearer_token_key(
headers, path, session, auth, max_cost_for_model
headers, path, session, auth, max_cost_for_model, model_id
)
else:
@@ -387,7 +447,12 @@ async def proxy(
async def get_bearer_token_key(
headers: dict, path: str, session: AsyncSession, auth: str, min_cost: int = 0
headers: dict,
path: str,
session: AsyncSession,
auth: str,
min_cost: int = 0,
model_id: str = "unknown",
) -> ApiKey:
"""Handle bearer token authentication proxy requests."""
parts = auth.split()
@@ -457,11 +522,13 @@ async def get_bearer_token_key(
except Exception as e:
key_preview = bearer_key[:20] + "..." if len(bearer_key) > 20 else bearer_key
logger.error(
f"Bearer token validation failed: {type(e).__name__}: {e} path={path} key={key_preview!r}",
f"Bearer token validation failed: {type(e).__name__}: {e} path={path} model={model_id!r} min_cost={min_cost} key={key_preview!r}",
extra={
"error": str(e),
"error_type": type(e).__name__,
"path": path,
"model_id": model_id,
"min_cost_msat": min_cost,
"bearer_key_preview": key_preview,
},
)

File diff suppressed because it is too large Load Diff

View File

@@ -3,7 +3,7 @@ from __future__ import annotations
import asyncio
import os
import re
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Callable
if TYPE_CHECKING:
from ..core.settings import Settings
@@ -122,12 +122,16 @@ async def get_all_models_with_overrides(
async def refresh_upstreams_models_periodically(
upstreams: list[BaseUpstreamProvider],
upstreams_provider: (
Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider]
),
) -> None:
"""Background task to periodically refresh models cache for all providers.
Args:
upstreams: List of upstream provider instances
upstreams_provider: Either a callable returning the live upstream list
(preferred — picks up providers added/changed via reinitialize_upstreams),
or a static list (legacy, will go stale after reinitialize_upstreams).
"""
import asyncio
import random
@@ -139,9 +143,14 @@ async def refresh_upstreams_models_periodically(
logger.info("Provider models refresh disabled (interval <= 0)")
return
def _resolve_upstreams() -> list[BaseUpstreamProvider]:
if callable(upstreams_provider):
return upstreams_provider()
return upstreams_provider
while True:
try:
for upstream in upstreams:
for upstream in _resolve_upstreams():
try:
await upstream.refresh_models_cache()
except Exception as e:

View File

@@ -65,9 +65,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
"""Strip 'ollama/' prefix for Ollama API compatibility."""
return model_id.removeprefix("ollama/")
def get_request_base_url(
self, path: str, model_obj: Model | None = None
) -> str:
def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str:
"""Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
return f"{self.base_url.rstrip('/')}/v1"
@@ -166,103 +164,3 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
},
)
return []
async def refresh_models_cache(self) -> None:
"""Refresh the in-memory models cache from upstream API."""
try:
from ..payment.models import _update_model_sats_pricing
from ..payment.price import sats_usd_price
models = await self.fetch_models()
models_with_fees = [self._apply_provider_fee_to_model(m) for m in models]
try:
sats_to_usd = sats_usd_price()
self._models_cache = [
_update_model_sats_pricing(m, sats_to_usd) for m in models_with_fees
]
except Exception:
self._models_cache = models_with_fees
self._models_by_id = {m.id: m for m in self._models_cache}
logger.info(
f"Refreshed models cache for {self.base_url}",
extra={"model_count": len(models)},
)
except Exception as e:
logger.error(
f"Failed to refresh models cache for {self.base_url}",
extra={"error": str(e), "error_type": type(e).__name__},
)
def get_cached_models(self) -> list[Model]:
"""Get cached models for this provider.
Returns:
List of cached Model objects
"""
return self._models_cache
def get_cached_model_by_id(self, model_id: str) -> Model | None:
"""Get a specific cached model by ID.
Args:
model_id: Model identifier
Returns:
Model object or None if not found
"""
return self._models_by_id.get(model_id)
def _apply_provider_fee_to_model(self, model: Model) -> Model:
"""Apply provider fee to model's USD pricing and calculate max costs.
Args:
model: Model object to update
Returns:
Model with provider fee applied to pricing and max costs calculated
"""
from ..payment.models import Model, Pricing, _calculate_usd_max_costs
adjusted_pricing = Pricing.parse_obj(
{k: v * self.provider_fee for k, v in model.pricing.dict().items()}
)
temp_model = Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=adjusted_pricing,
sats_pricing=None,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
)
(
adjusted_pricing.max_prompt_cost,
adjusted_pricing.max_completion_cost,
adjusted_pricing.max_cost,
) = _calculate_usd_max_costs(temp_model)
return Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=adjusted_pricing,
sats_pricing=model.sats_pricing,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
)

View File

@@ -1,9 +1,9 @@
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Optional
import httpx
from pydantic import BaseModel
from pydantic import BaseModel, Field
from ..core.logging import get_logger
from ..payment.models import Architecture, Model, Pricing, async_fetch_openrouter_models
@@ -16,18 +16,20 @@ logger = get_logger(__name__)
class PPQAIModelPricing(BaseModel):
ui: dict[str, float]
api: dict[str, float]
ui: Optional[dict[str, float]] = None
api: Optional[dict[str, float]] = None
input_per_1M_tokens: Optional[float] = Field(None, alias="input_per_1M_tokens")
output_per_1M_tokens: Optional[float] = Field(None, alias="output_per_1M_tokens")
class PPQAIModel(BaseModel):
id: str
provider: str
provider: Optional[str] = None
name: str
created_at: int
context_length: int
pricing: PPQAIModelPricing
popular: bool
popular: bool = False
class PPQAIUpstreamProvider(BaseUpstreamProvider):
@@ -134,31 +136,54 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
)
if or_model:
if input_price := ppqai_model.pricing.api.get(
"input_per_1M"
):
input_price = None
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M"
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
if input_price is not None:
or_model.pricing.prompt = input_price / 1_000_000
if output_price := ppqai_model.pricing.api.get(
"output_per_1M"
):
output_price = None
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M"
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
if output_price is not None:
or_model.pricing.completion = output_price / 1_000_000
if cl := ppqai_model.context_length:
or_model.context_length = cl
models.append(or_model)
else:
input_price = ppqai_model.pricing.api.get(
"input_per_1M", 0.0
)
output_price = ppqai_model.pricing.api.get(
"output_per_1M", 0.0
)
input_price = 0.0
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M", 0.0
)
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
output_price = 0.0
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M", 0.0
)
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
models.append(
Model(
id=ppqai_model.id,
name=ppqai_model.name,
created=ppqai_model.created_at // 1000,
description=f"{ppqai_model.provider} model",
description=f"{ppqai_model.provider or 'PPQ.AI'} model",
context_length=ppqai_model.context_length,
architecture=Architecture(
modality="text->text",

View File

@@ -8,6 +8,7 @@ from cashu.wallet.wallet import Wallet
from sqlmodel import col, select, update
from .core import db, get_logger
from .core.db import store_cashu_transaction
from .core.settings import settings
from .payment.lnurl import raw_send_to_lnurl
@@ -155,6 +156,20 @@ async def swap_to_primary_mint(
raise ValueError("Invalid unit")
primary_wallet = await get_wallet(settings.primary_mint, settings.primary_mint_unit)
# If the token is already from the primary mint, we don't need to swap
# and we definitely don't want to calculate or pay fees.
if token_obj.mint == settings.primary_mint:
logger.info(
"swap_to_primary_mint: token already on primary mint, skipping swap",
extra={
"mint": token_obj.mint,
"amount": token_amount,
"unit": token_obj.unit,
},
)
await token_wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True)
return token_amount, token_obj.unit, token_obj.mint
minted_amount = await _calculate_swap_amount(
amount_msat,
token_obj.unit,
@@ -265,6 +280,8 @@ async def credit_balance(
try:
amount, unit, mint_url = await recieve_token(cashu_token)
original_amount = amount
original_unit = unit
logger.info(
"credit_balance: Token redeemed successfully",
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
@@ -296,6 +313,19 @@ async def credit_balance(
extra={"new_balance": key.balance},
)
try:
await store_cashu_transaction(
token=cashu_token,
amount=original_amount,
unit=original_unit,
mint_url=mint_url,
typ="in",
source="apikey",
api_key_hashed_key=key.hashed_key,
)
except Exception:
pass
logger.info(
"Cashu token successfully redeemed and stored",
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
@@ -535,7 +565,7 @@ async def periodic_refund_sweep() -> None:
except Exception as e:
error_msg = str(e).lower()
if "already spent" in error_msg:
refund.swept = True
refund.collected = True
session.add(refund)
logger.info(
"Refund already spent (client collected), marking swept",
@@ -559,6 +589,46 @@ async def periodic_refund_sweep() -> None:
)
async def periodic_routstr_fee_payout() -> None:
from .auth import (
ROUTSTR_FEE_DEFAULT_PAYOUT,
ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS,
ROUTSTR_LN_ADDRESS,
)
if not ROUTSTR_LN_ADDRESS:
logger.info("ROUTSTR_LN_ADDRESS not set, skipping fee payout")
return
while True:
await asyncio.sleep(ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS)
try:
async with db.create_session() as session:
fee = await db.get_routstr_fee(session)
accumulated_sats = fee.accumulated_msats // 1000
if accumulated_sats >= ROUTSTR_FEE_DEFAULT_PAYOUT:
wallet = await get_wallet(settings.primary_mint, "sat")
proofs = get_proofs_per_mint_and_unit(
wallet, settings.primary_mint, "sat", not_reserved=True
)
amount_received = await raw_send_to_lnurl(
wallet, proofs, ROUTSTR_LN_ADDRESS, "sat", amount=accumulated_sats
)
paid_msats = accumulated_sats * 1000
await db.reset_routstr_fee(session, paid_msats)
logger.info(
"Routstr fee payout sent",
extra={
"accumulated_sats": accumulated_sats,
"amount_received": amount_received,
},
)
except Exception as e:
logger.error(
f"Error in Routstr fee payout: {type(e).__name__}",
extra={"error": str(e)},
)
async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int:
wallet = await get_wallet(mint, unit)
proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id]

View File

@@ -380,6 +380,13 @@ async def integration_session(
yield session
@pytest_asyncio.fixture
async def patched_db_engine(integration_engine: Any) -> AsyncGenerator[None, None]:
"""Patch the global db engine so create_session() uses the test engine."""
with patch("routstr.core.db.engine", integration_engine):
yield
class DatabaseSnapshot:
"""Utility to capture and compare database states"""

View File

@@ -0,0 +1,369 @@
"""
Integration tests for the balance-goes-negative bug in adjust_payment_for_tokens.
Root cause: when actual token cost exceeds the discounted reservation
(cost_difference > 0, caused by tolerance_percentage discounting the reservation),
the finalization UPDATE had no WHERE guard on balance, allowing balance to go negative.
Fix: added `.where(col(ApiKey.balance) >= total_cost_msats)` so the UPDATE is a no-op
when balance is insufficient, then falls back to charging only deducted_max_cost.
"""
import uuid
from unittest.mock import patch
import pytest
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey
from routstr.payment.cost_calculation import CostData
def _make_key(balance: int, reserved: int) -> ApiKey:
return ApiKey(
hashed_key=f"test_{uuid.uuid4().hex}",
balance=balance,
reserved_balance=reserved,
total_spent=0,
total_requests=1,
)
async def _refresh(session: AsyncSession, key: ApiKey) -> ApiKey:
await session.refresh(key)
return key
# ---------------------------------------------------------------------------
# Helper: build a CostData where token cost > deducted_max_cost
# ---------------------------------------------------------------------------
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,
)
# ---------------------------------------------------------------------------
# Test 1 — exact reproduction of the bug
#
# Setup: balance == deducted_max_cost (user has just enough for the reservation,
# nothing extra). Actual token cost is 1% higher (tolerance_percentage).
#
# Before fix: balance -= total_cost_msats → goes negative.
# After fix: WHERE balance >= total_cost_msats fails → fallback charges
# deducted_max_cost → balance reaches 0, never negative.
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_balance_never_negative_when_cost_exceeds_reservation(
integration_session: AsyncSession,
) -> None:
"""Balance must not go negative when actual token cost > discounted reservation."""
from routstr.auth import adjust_payment_for_tokens
deducted_max_cost = 990 # reserved (1% below true max of 1000)
actual_token_cost = 1000 # actual cost at true max
# User has balance exactly equal to the reservation — tight budget
key = _make_key(balance=deducted_max_cost, reserved=deducted_max_cost)
integration_session.add(key)
await integration_session.commit()
response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}}
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await _refresh(integration_session, key)
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.reserved_balance == 0, "Reservation must be fully released after finalization"
# ---------------------------------------------------------------------------
# Test 2 — balance is ZERO after the reservation is accounted for
# (absolute floor case)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_balance_floor_at_zero_on_overrun(
integration_session: AsyncSession,
) -> None:
"""When balance exactly covers deducted_max_cost and cost overruns, balance reaches 0 not negative."""
from routstr.auth import adjust_payment_for_tokens
deducted_max_cost = 500
actual_token_cost = 550 # 10% overrun
key = _make_key(balance=500, reserved=500)
integration_session.add(key)
await integration_session.commit()
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await _refresh(integration_session, key)
assert key.balance == 0, (
f"Expected balance=0 (charged deducted_max_cost fallback), got {key.balance}"
)
assert key.reserved_balance == 0, f"Reserved balance should be 0, got {key.reserved_balance}"
# Fallback charges deducted_max_cost
assert key.total_spent == deducted_max_cost, (
f"Expected total_spent={deducted_max_cost}, got {key.total_spent}"
)
# ---------------------------------------------------------------------------
# Test 3 — balance has enough room: full token cost should be charged
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_full_cost_charged_when_balance_sufficient_for_overrun(
integration_session: AsyncSession,
) -> None:
"""When balance covers total_cost_msats, the full amount is charged (not just deducted_max_cost)."""
from routstr.auth import adjust_payment_for_tokens
deducted_max_cost = 990
actual_token_cost = 1000
# User has extra balance beyond the reservation
key = _make_key(balance=2000, reserved=990)
integration_session.add(key)
await integration_session.commit()
response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}}
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await _refresh(integration_session, key)
assert key.balance >= 0, f"Balance went negative: {key.balance}"
assert key.reserved_balance == 0, f"Reservation not released: {key.reserved_balance}"
assert key.total_spent == actual_token_cost, (
f"Expected full charge of {actual_token_cost}, got {key.total_spent}"
)
assert key.balance == 2000 - actual_token_cost, (
f"Expected balance={2000 - actual_token_cost}, got {key.balance}"
)
# ---------------------------------------------------------------------------
# Test 4 — concurrent finalizations with cost overrun
#
# Multiple requests finish concurrently. Each has a small overrun.
# None should drive balance negative.
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_concurrent_cost_overruns_never_negative(
integration_session: AsyncSession,
patched_db_engine: None,
) -> None:
"""Concurrent finalization with cost overruns must never produce negative balance."""
import asyncio
from routstr.auth import adjust_payment_for_tokens, pay_for_request
from routstr.core.db import create_session
deducted_max_cost = 990
actual_token_cost = 1000
n_requests = 5
# Fund the key with exactly enough for n_requests reservations + a tiny buffer
starting_balance = deducted_max_cost * n_requests
key_hash = f"test_concurrent_{uuid.uuid4().hex}"
async with create_session() as session:
key = ApiKey(
hashed_key=key_hash,
balance=starting_balance,
reserved_balance=0,
total_spent=0,
total_requests=0,
)
session.add(key)
await session.commit()
# Reserve n_requests slots (sequentially, as pay_for_request is atomic)
async with create_session() as session:
key_to_reserve = await session.get(ApiKey, key_hash)
assert key_to_reserve is not None
for _ in range(n_requests):
await pay_for_request(key_to_reserve, deducted_max_cost, session)
await session.refresh(key_to_reserve)
# Now finalize all concurrently with cost overrun
async def finalize() -> None:
response_data = {
"model": "test-model",
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
}
async with create_session() as session:
fresh_key = await session.get(ApiKey, key_hash)
assert fresh_key is not None
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
fresh_key, response_data, session, deducted_max_cost
)
await asyncio.gather(*[finalize() for _ in range(n_requests)])
async with create_session() as session:
final_key = await session.get(ApiKey, key_hash)
assert final_key is not None
assert final_key.balance >= 0, (
f"Balance went negative after concurrent overruns: {final_key.balance}"
)
assert final_key.reserved_balance == 0, (
f"Reserved balance not fully released: {final_key.reserved_balance}"
)
assert final_key.total_spent <= starting_balance, (
f"Total spent ({final_key.total_spent}) exceeds starting balance ({starting_balance})"
)
# Every request must have been charged at least deducted_max_cost — no free inference.
assert final_key.total_spent == starting_balance, (
f"Expected total_spent={starting_balance} (all {n_requests} reservations charged), "
f"got {final_key.total_spent} — at least one request got free inference"
)
# ---------------------------------------------------------------------------
# Test 5 — overrun with no balance at all (reserved_balance == balance)
# simulates a user who topped up to exactly the reservation floor
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_zero_free_balance_overrun_is_safe(
integration_session: AsyncSession,
) -> None:
"""User with zero free balance (all reserved) should never go negative on overrun."""
from routstr.auth import adjust_payment_for_tokens
deducted_max_cost = 1000
actual_token_cost = 1050
# balance == reserved_balance: zero free balance
key = _make_key(balance=1000, reserved=1000)
integration_session.add(key)
await integration_session.commit()
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 100}}
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
await _refresh(integration_session, key)
assert key.balance >= 0, f"Balance went negative: {key.balance}"
assert key.reserved_balance >= 0, f"Reserved balance went negative: {key.reserved_balance}"
# ---------------------------------------------------------------------------
# Test 6 — parallel requests: second finalization must not get free inference
#
# Root cause of the bug fixed in auth.py:
# `.where(col(ApiKey.balance) >= total_cost_msats)` ignores other requests'
# reservations, so after Request A charges total_cost_msats, balance can drop
# below deducted_max_cost, causing Request B's fallback to release for free.
#
# Fix: use `balance - reserved_balance + deducted_max_cost >= total_cost_msats`
# so the check accounts for concurrent reservations.
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_parallel_requests_no_free_inference(
integration_session: AsyncSession,
patched_db_engine: None,
) -> None:
"""Second parallel finalization must be charged even when first depleted free balance."""
import asyncio
from routstr.auth import adjust_payment_for_tokens
from routstr.core.db import create_session
deducted_max_cost = 100
actual_token_cost = 150 # overrun: 50 more than reserved
# Fund the key with exactly 2 * deducted_max_cost.
# Both requests pre-reserved 100 each → balance=200, reserved=200, free=0.
# Old check (balance >= total_cost_msats):
# Request A: 200 >= 150 ✓ → charges 150 → balance=50, reserved=100
# Request B: 50 >= 150 ✗ → fallback: 50 >= 100 ✗ → releases FREE
# New check (balance - reserved + deducted >= total_cost_msats):
# Both fall to fallback (0 free balance).
# Both charge deducted_max_cost=100 → total_spent=200, balance=0.
starting_balance = deducted_max_cost * 2
key_hash = f"test_parallel_no_free_{uuid.uuid4().hex}"
async with create_session() as session:
key = ApiKey(
hashed_key=key_hash,
balance=starting_balance,
reserved_balance=deducted_max_cost * 2, # both slots pre-reserved
total_spent=0,
total_requests=2,
)
session.add(key)
await session.commit()
async def finalize() -> None:
response_data = {
"model": "test-model",
"usage": {"prompt_tokens": 50, "completion_tokens": 100},
}
async with create_session() as session:
fresh_key = await session.get(ApiKey, key_hash)
assert fresh_key is not None
with patch(
"routstr.auth.calculate_cost",
return_value=_cost_data(actual_token_cost),
):
await adjust_payment_for_tokens(
fresh_key, response_data, session, deducted_max_cost
)
await asyncio.gather(finalize(), finalize())
async with create_session() as session:
final_key = await session.get(ApiKey, key_hash)
assert final_key is not None
assert final_key.balance >= 0, f"Balance went negative: {final_key.balance}"
assert final_key.reserved_balance == 0, (
f"Reserved balance not released: {final_key.reserved_balance}"
)
# Both requests must have been charged — no free inference.
assert final_key.total_spent == starting_balance, (
f"Expected total_spent={starting_balance} (both reservations charged), "
f"got {final_key.total_spent} — one request got free inference"
)

View File

@@ -0,0 +1,302 @@
"""Integration tests for CLI token management (/admin/api/cli-tokens).
Covers:
- GET /admin/api/cli-tokens — list (preview only, no full token)
- POST /admin/api/cli-tokens — create (returns full token once)
- DELETE /admin/api/cli-tokens/{id} — revoke
- Using a CLI token as Bearer auth against admin endpoints
- Expiry enforcement (expired tokens are rejected by require_admin_api)
- last_used_at bump on successful use
- Auth failures: missing token, wrong token, revoked token
"""
from __future__ import annotations
import secrets
import time
from typing import AsyncGenerator
import pytest
import pytest_asyncio
from httpx import AsyncClient
from sqlmodel import select
from routstr.core.admin import admin_sessions
from routstr.core.db import AsyncSession, CliToken
# ──────────────────────────────────────────────────────────────────────────────
# Fixtures
# ──────────────────────────────────────────────────────────────────────────────
@pytest_asyncio.fixture
async def admin_session_token() -> AsyncGenerator[str, None]:
"""Inject a short-lived admin session token into admin_sessions."""
token = secrets.token_urlsafe(24)
admin_sessions[token] = int(time.time()) + 3600
yield token
admin_sessions.pop(token, None)
@pytest_asyncio.fixture
async def admin_client(
integration_client: AsyncClient, admin_session_token: str
) -> AsyncClient:
"""An integration_client pre-authenticated with an admin session token."""
integration_client.headers["Authorization"] = f"Bearer {admin_session_token}"
return integration_client
# ──────────────────────────────────────────────────────────────────────────────
# Creation
# ──────────────────────────────────────────────────────────────────────────────
@pytest.mark.integration
@pytest.mark.asyncio
async def test_create_cli_token_returns_full_token_once(
admin_client: AsyncClient,
) -> None:
"""POST /admin/api/cli-tokens returns the raw token only on creation."""
resp = await admin_client.post(
"/admin/api/cli-tokens",
json={"name": "my-laptop"},
)
assert resp.status_code == 200
body = resp.json()
assert body["name"] == "my-laptop"
assert isinstance(body["id"], str) and body["id"]
assert isinstance(body["token"], str) and len(body["token"]) >= 32
assert body["expires_at"] is None
assert isinstance(body["created_at"], int)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_create_cli_token_with_expiry(admin_client: AsyncClient) -> None:
"""expires_in_days sets expires_at ~= now + days * 86400."""
before = int(time.time())
resp = await admin_client.post(
"/admin/api/cli-tokens",
json={"name": "ci-runner", "expires_in_days": 7},
)
assert resp.status_code == 200
body = resp.json()
assert body["expires_at"] is not None
delta = body["expires_at"] - before
# Allow 10s jitter around 7 * 86400
assert 7 * 86400 - 10 <= delta <= 7 * 86400 + 10
@pytest.mark.integration
@pytest.mark.asyncio
async def test_create_cli_token_rejects_empty_name(
admin_client: AsyncClient,
) -> None:
resp = await admin_client.post(
"/admin/api/cli-tokens", json={"name": " "}
)
assert resp.status_code == 400
@pytest.mark.integration
@pytest.mark.asyncio
async def test_create_cli_token_requires_admin(
integration_client: AsyncClient,
) -> None:
"""No admin token / no bearer → 403."""
resp = await integration_client.post(
"/admin/api/cli-tokens", json={"name": "no-auth"}
)
assert resp.status_code == 403
# ──────────────────────────────────────────────────────────────────────────────
# Listing
# ──────────────────────────────────────────────────────────────────────────────
@pytest.mark.integration
@pytest.mark.asyncio
async def test_list_cli_tokens_returns_preview_not_full_token(
admin_client: AsyncClient,
) -> None:
"""Listing never leaks the raw token."""
create = await admin_client.post(
"/admin/api/cli-tokens", json={"name": "secret-keeper"}
)
assert create.status_code == 200
full_token = create.json()["token"]
resp = await admin_client.get("/admin/api/cli-tokens")
assert resp.status_code == 200
items = resp.json()
assert any(t["name"] == "secret-keeper" for t in items)
for t in items:
# No 'token' field, only 'token_preview'
assert "token" not in t
assert "token_preview" in t
assert full_token not in t["token_preview"]
assert "..." in t["token_preview"]
@pytest.mark.integration
@pytest.mark.asyncio
async def test_list_cli_tokens_requires_admin(
integration_client: AsyncClient,
) -> None:
resp = await integration_client.get("/admin/api/cli-tokens")
assert resp.status_code == 403
# ──────────────────────────────────────────────────────────────────────────────
# Using a CLI token as admin auth
# ──────────────────────────────────────────────────────────────────────────────
@pytest.mark.integration
@pytest.mark.asyncio
async def test_cli_token_authorizes_admin_endpoints(
admin_client: AsyncClient,
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A freshly-created CLI token can be used as Bearer on admin endpoints."""
create = await admin_client.post(
"/admin/api/cli-tokens", json={"name": "cli-auth"}
)
assert create.status_code == 200
cli_token = create.json()["token"]
token_id = create.json()["id"]
# Use a NEW client to isolate the header from admin_session_token
integration_client.headers["Authorization"] = f"Bearer {cli_token}"
resp = await integration_client.get("/admin/api/cli-tokens")
assert resp.status_code == 200
# last_used_at should be populated after use
row = await integration_session.get(CliToken, token_id)
assert row is not None
assert row.last_used_at is not None
assert row.last_used_at >= row.created_at
@pytest.mark.integration
@pytest.mark.asyncio
async def test_expired_cli_token_is_rejected(
admin_client: AsyncClient,
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""A CLI token with expires_at in the past → 403."""
create = await admin_client.post(
"/admin/api/cli-tokens",
json={"name": "will-expire", "expires_in_days": 1},
)
assert create.status_code == 200
cli_token = create.json()["token"]
token_id = create.json()["id"]
# Force-expire it in the DB
row = await integration_session.get(CliToken, token_id)
assert row is not None
row.expires_at = int(time.time()) - 1
integration_session.add(row)
await integration_session.commit()
integration_client.headers["Authorization"] = f"Bearer {cli_token}"
resp = await integration_client.get("/admin/api/cli-tokens")
assert resp.status_code == 403
@pytest.mark.integration
@pytest.mark.asyncio
async def test_invalid_bearer_token_is_rejected(
integration_client: AsyncClient,
) -> None:
integration_client.headers["Authorization"] = "Bearer not-a-real-token"
resp = await integration_client.get("/admin/api/cli-tokens")
assert resp.status_code == 403
# ──────────────────────────────────────────────────────────────────────────────
# Revocation
# ──────────────────────────────────────────────────────────────────────────────
@pytest.mark.integration
@pytest.mark.asyncio
async def test_revoke_cli_token_removes_auth(
admin_client: AsyncClient,
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""After DELETE, the token no longer authorizes."""
create = await admin_client.post(
"/admin/api/cli-tokens", json={"name": "to-revoke"}
)
token_id = create.json()["id"]
cli_token = create.json()["token"]
revoke = await admin_client.delete(f"/admin/api/cli-tokens/{token_id}")
assert revoke.status_code == 200
assert revoke.json() == {"ok": True, "deleted_id": token_id}
# Row is gone
row = await integration_session.get(CliToken, token_id)
assert row is None
# Can no longer be used for auth
integration_client.headers["Authorization"] = f"Bearer {cli_token}"
resp = await integration_client.get("/admin/api/cli-tokens")
assert resp.status_code == 403
@pytest.mark.integration
@pytest.mark.asyncio
async def test_revoke_unknown_cli_token_returns_404(
admin_client: AsyncClient,
) -> None:
resp = await admin_client.delete("/admin/api/cli-tokens/does-not-exist")
assert resp.status_code == 404
@pytest.mark.integration
@pytest.mark.asyncio
async def test_revoke_cli_token_requires_admin(
integration_client: AsyncClient,
) -> None:
resp = await integration_client.delete("/admin/api/cli-tokens/anything")
assert resp.status_code == 403
# ──────────────────────────────────────────────────────────────────────────────
# Lifecycle / uniqueness
# ──────────────────────────────────────────────────────────────────────────────
@pytest.mark.integration
@pytest.mark.asyncio
async def test_multiple_tokens_are_independent(
admin_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""Creating N tokens yields N unique tokens that all live in DB."""
names = ["dev-a", "dev-b", "dev-c"]
raw_tokens: list[str] = []
ids: list[str] = []
for name in names:
r = await admin_client.post(
"/admin/api/cli-tokens", json={"name": name}
)
assert r.status_code == 200
raw_tokens.append(r.json()["token"])
ids.append(r.json()["id"])
# All unique
assert len(set(raw_tokens)) == len(raw_tokens)
assert len(set(ids)) == len(ids)
# All in DB
result = await integration_session.exec(
select(CliToken).where(CliToken.name.in_(names)) # type: ignore[attr-defined]
)
rows = result.all()
assert {r.name for r in rows} == set(names)

View File

@@ -0,0 +1,240 @@
"""
Tests showing how a user hits "Insufficient balance: X mSats required for this model"
when their balance is too low for the model's cost.
The log line that triggered this:
WARNING Insufficient billing balance during validation
ERROR Bearer token validation failed: HTTPException: 402:
{'error': {'message': 'Insufficient balance: 622888 mSats required
for this model. 20320 available.', ...}}
This happens in validate_bearer_key (auth.py) when:
billing_key.total_balance < min_cost (model's max cost)
and also in pay_for_request when the atomic UPDATE finds no available balance.
"""
import uuid
from unittest.mock import patch
import pytest
from fastapi import HTTPException
from httpx import AsyncClient
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey
def _key(balance: int, reserved: int = 0) -> ApiKey:
return ApiKey(
hashed_key=f"test_{uuid.uuid4().hex}",
balance=balance,
reserved_balance=reserved,
total_spent=0,
)
# ---------------------------------------------------------------------------
# Test 1 — simplest case: balance < model cost → pay_for_request raises 402
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pay_for_request_raises_402_when_balance_too_low(
integration_session: AsyncSession,
) -> None:
"""
User has 20_000 msats. Model costs 622_888 msats.
pay_for_request must raise HTTP 402 with a clear message.
"""
from routstr.auth import pay_for_request
model_cost = 622_888
user_balance = 20_000
key = _key(balance=user_balance)
integration_session.add(key)
await integration_session.commit()
with pytest.raises(HTTPException) as exc_info:
await pay_for_request(key, model_cost, integration_session)
assert exc_info.value.status_code == 402
detail = exc_info.value.detail
assert isinstance(detail, dict)
error = detail["error"]
assert error["code"] == "insufficient_balance"
assert str(model_cost) in error["message"]
assert str(user_balance) in error["message"]
# Balance must be untouched
await integration_session.refresh(key)
assert key.balance == user_balance
assert key.reserved_balance == 0
# ---------------------------------------------------------------------------
# Test 2 — balance is zero
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pay_for_request_raises_402_on_zero_balance(
integration_session: AsyncSession,
) -> None:
"""User with zero balance cannot make any request."""
from routstr.auth import pay_for_request
key = _key(balance=0)
integration_session.add(key)
await integration_session.commit()
with pytest.raises(HTTPException) as exc_info:
await pay_for_request(key, 1_000, integration_session)
assert exc_info.value.status_code == 402
detail = exc_info.value.detail
assert isinstance(detail, dict)
assert detail["error"]["code"] == "insufficient_balance"
# ---------------------------------------------------------------------------
# Test 3 — all balance is reserved (total_balance = balance - reserved = 0)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pay_for_request_raises_402_when_all_balance_reserved(
integration_session: AsyncSession,
) -> None:
"""
User has 50_000 msats balance but 50_000 is already reserved for in-flight
requests. Free balance (total_balance) = 0. Should get 402.
"""
from routstr.auth import pay_for_request
key = _key(balance=50_000, reserved=50_000)
integration_session.add(key)
await integration_session.commit()
with pytest.raises(HTTPException) as exc_info:
await pay_for_request(key, 1_000, integration_session)
assert exc_info.value.status_code == 402
# Balance and reserved must be untouched
await integration_session.refresh(key)
assert key.balance == 50_000
assert key.reserved_balance == 50_000
# ---------------------------------------------------------------------------
# Test 4 — balance just one msat below model cost
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pay_for_request_raises_402_one_msat_short(
integration_session: AsyncSession,
) -> None:
"""Off-by-one: balance is exactly model_cost - 1."""
from routstr.auth import pay_for_request
model_cost = 10_000
key = _key(balance=model_cost - 1)
integration_session.add(key)
await integration_session.commit()
with pytest.raises(HTTPException) as exc_info:
await pay_for_request(key, model_cost, integration_session)
assert exc_info.value.status_code == 402
await integration_session.refresh(key)
assert key.balance == model_cost - 1 # untouched
assert key.reserved_balance == 0
# ---------------------------------------------------------------------------
# Test 5 — balance exactly equal to model cost → succeeds
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pay_for_request_succeeds_when_balance_equals_cost(
integration_session: AsyncSession,
) -> None:
"""Balance == model cost: the request should be reserved successfully."""
from routstr.auth import pay_for_request
model_cost = 10_000
key = _key(balance=model_cost)
integration_session.add(key)
await integration_session.commit()
# Should not raise
await pay_for_request(key, model_cost, integration_session)
await integration_session.refresh(key)
assert key.reserved_balance == model_cost
assert key.balance == model_cost # balance unchanged, only reserved goes up
# ---------------------------------------------------------------------------
# Test 6 — HTTP layer returns 402 JSON with the right shape
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_http_402_response_shape_on_insufficient_balance(
integration_client: AsyncClient,
integration_session: AsyncSession,
) -> None:
"""
End-to-end: POST /v1/chat/completions with a key whose balance is far below
the mocked model cost returns HTTP 402 with the expected JSON error body.
Matches exactly the log snippet in the bug report:
'Insufficient balance: X mSats required for this model. Y available.'
"""
from unittest.mock import AsyncMock, MagicMock
model_cost = 622_888
user_balance = 20_320
key = _key(balance=user_balance)
integration_session.add(key)
await integration_session.commit()
# Minimal model stub so proxy routing doesn't 400 before reaching balance check
mock_model = MagicMock()
mock_model.sats_pricing = None
# Upstream stub — never reached because balance check fires first
mock_upstream = MagicMock()
mock_upstream.prepare_headers = MagicMock(return_value={})
with (
patch("routstr.proxy.get_model_instance", return_value=mock_model),
patch("routstr.proxy.get_provider_for_model", return_value=[mock_upstream]),
# Patch where it is used (proxy imports it at module level)
patch(
"routstr.proxy.get_max_cost_for_model",
new=AsyncMock(return_value=model_cost),
),
):
response = await integration_client.post(
"/v1/chat/completions",
headers={"Authorization": f"Bearer sk-{key.hashed_key}"},
json={
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hello"}],
},
)
assert response.status_code == 402
body = response.json()
# FastAPI wraps HTTPException detail under "detail"
error = body["detail"]["error"]
assert error["code"] == "insufficient_balance"
assert error["type"] == "insufficient_quota"
assert str(model_cost) in error["message"]
assert str(user_balance) in error["message"]
# Balance must be completely untouched
await integration_session.refresh(key)
assert key.balance == user_balance
assert key.reserved_balance == 0
assert key.total_spent == 0

View File

@@ -0,0 +1,190 @@
"""Integration tests for model price updates when provider fee schedules change."""
import time
from typing import Any, Generator
import pytest
from httpx import AsyncClient
from routstr.core.admin import admin_sessions
ADMIN_TOKEN = "test-admin-token"
def _auth_header() -> dict[str, str]:
return {"Authorization": f"Bearer {ADMIN_TOKEN}"}
@pytest.fixture(autouse=True)
def _inject_admin_session() -> Generator[None, None, None]:
admin_sessions[ADMIN_TOKEN] = int(time.time()) + 3600
yield
admin_sessions.pop(ADMIN_TOKEN, None)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_model_price_updates_on_fee_schedule_change(
integration_client: AsyncClient,
patched_db_engine: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Patch fetch_models to return empty list to avoid network errors
# and allow DB models to be used
from routstr.upstream.base import BaseUpstreamProvider
async def mock_fetch_models(self: BaseUpstreamProvider) -> list:
return []
monkeypatch.setattr(BaseUpstreamProvider, "fetch_models", mock_fetch_models)
# 1. Create a provider
provider_resp = await integration_client.post(
"/admin/api/upstream-providers",
json={
"provider_type": "custom",
"base_url": "https://api.example.com/v1",
"api_key": "test-key",
"enabled": True,
"provider_fee": 1.0,
},
headers=_auth_header(),
)
provider_id = provider_resp.json()["id"]
# 2. Add a model to this provider
model_id = "test-model-price-update"
await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/models",
json={
"id": model_id,
"name": "Test Model",
"created": int(time.time()),
"description": "Test",
"context_length": 4096,
"architecture": {
"modality": "text",
"input_modalities": ["text"],
"output_modalities": ["text"],
"tokenizer": "gpt2",
"instruct_type": "none",
},
"pricing": {"prompt": 1.0, "completion": 2.0},
"enabled": True,
},
headers=_auth_header(),
)
# 3. Check initial price (should be prompt=1.0 * fee=1.0 = 1.0)
# We use /models endpoint
resp = await integration_client.get("/models")
models = resp.json()["data"]
target = next((m for m in models if m["id"] == model_id), None)
assert target is not None
assert target["pricing"]["prompt"] == 1.0
# 4. Update provider fee schedule to a very high value for the current time
# We'll use a range that covers the whole day to be safe
schedules = [
{"start_time": "00:00", "end_time": "23:59", "provider_fee": 2.5},
]
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={"schedules": schedules},
headers=_auth_header(),
)
# 5. Check price again - should be updated instantly
resp = await integration_client.get("/models")
models = resp.json()["data"]
target = next((m for m in models if m["id"] == model_id), None)
assert target is not None
# 1.0 * 2.5 = 2.5
assert target["pricing"]["prompt"] == 2.5
# 6. Delete schedules
await integration_client.delete(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
headers=_auth_header(),
)
# 7. Should revert to default fee (1.0)
resp = await integration_client.get("/models")
models = resp.json()["data"]
target = next((m for m in models if m["id"] == model_id), None)
assert target is not None
assert target["pricing"]["prompt"] == 1.0
@pytest.mark.integration
@pytest.mark.asyncio
async def test_upstream_model_price_updates_on_fee_schedule_change(
integration_client: AsyncClient,
patched_db_engine: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
from routstr.payment.models import Architecture, Model, Pricing
from routstr.upstream.base import BaseUpstreamProvider
upstream_model_id = "upstream-model-only"
# Mock fetch_models to return a model
async def mock_fetch_models(self: BaseUpstreamProvider) -> list[Model]:
return [
Model(
id=upstream_model_id,
name="Upstream Model",
created=int(time.time()),
description="Test",
context_length=4096,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="gpt2",
instruct_type="none",
),
pricing=Pricing(prompt=1.0, completion=2.0),
enabled=True,
)
]
monkeypatch.setattr(BaseUpstreamProvider, "fetch_models", mock_fetch_models)
# 1. Create a provider
provider_resp = await integration_client.post(
"/admin/api/upstream-providers",
json={
"provider_type": "custom",
"base_url": "https://api.example.com/v1",
"api_key": "test-key-2",
"enabled": True,
"provider_fee": 1.0,
},
headers=_auth_header(),
)
provider_id = provider_resp.json()["id"]
# 2. Check initial price (should be prompt=1.0 * fee=1.0 = 1.0)
resp = await integration_client.get("/models")
models = resp.json()["data"]
target = next((m for m in models if m["id"] == upstream_model_id), None)
assert target is not None
assert target["pricing"]["prompt"] == 1.0
# 3. Update provider fee schedule
schedules = [
{"start_time": "00:00", "end_time": "23:59", "provider_fee": 3.0},
]
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={"schedules": schedules},
headers=_auth_header(),
)
# 4. Check price again - I expect this to FAIL (still 1.0 instead of 3.0)
resp = await integration_client.get("/models")
models = resp.json()["data"]
target = next((m for m in models if m["id"] == upstream_model_id), None)
assert target is not None
assert target["pricing"]["prompt"] == 3.0

View File

@@ -101,7 +101,7 @@ async def test_enforce_lowest_provider_fee_for_same_url(
)
]
async def refresh_models_cache(self) -> None:
async def refresh_models_cache(self, skip_network: bool = False) -> None:
pass
def prepare_headers(self, request_headers: dict[str, str]) -> dict[str, str]:

View File

@@ -0,0 +1,402 @@
"""Integration tests for provider fee schedule API endpoints."""
from typing import Any, Generator
import pytest
from httpx import AsyncClient
from routstr.core.admin import admin_sessions
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
ADMIN_TOKEN = "test-admin-token"
def _auth_header() -> dict[str, str]:
return {"Authorization": f"Bearer {ADMIN_TOKEN}"}
async def _create_provider(client: AsyncClient, *, fee: float = 1.02) -> int:
"""Create a test provider and return its ID."""
resp = await client.post(
"/admin/api/upstream-providers",
json={
"provider_type": "custom",
"base_url": "https://api.example.com/v1",
"api_key": "test-key",
"enabled": True,
"provider_fee": fee,
},
headers=_auth_header(),
)
assert resp.status_code == 200, resp.text
return resp.json()["id"]
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture(autouse=True)
def _inject_admin_session() -> Generator[None, None, None]:
"""Inject a valid admin session token for all tests."""
import time
admin_sessions[ADMIN_TOKEN] = int(time.time()) + 3600
yield
admin_sessions.pop(ADMIN_TOKEN, None)
@pytest.fixture(autouse=True)
def _patch_reinitialize(monkeypatch: Any) -> None:
async def _noop(*args: Any, **kwargs: Any) -> None:
pass
monkeypatch.setattr("routstr.core.admin.reinitialize_upstreams", _noop)
monkeypatch.setattr("routstr.core.admin.refresh_model_maps", _noop)
# ---------------------------------------------------------------------------
# GET fee schedules
# ---------------------------------------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_get_fee_schedules_empty_for_new_provider(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
resp = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
headers=_auth_header(),
)
assert resp.status_code == 200
assert resp.json() == []
@pytest.mark.integration
@pytest.mark.asyncio
async def test_get_fee_schedules_404_for_missing_provider(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
resp = await integration_client.get(
"/admin/api/upstream-providers/99999/fee-schedules",
headers=_auth_header(),
)
assert resp.status_code == 404
# ---------------------------------------------------------------------------
# PUT fee schedules
# ---------------------------------------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_put_fee_schedules_success(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
schedules = [
{"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05},
{"start_time": "18:00", "end_time": "08:00", "provider_fee": 1.02},
]
resp = await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={"schedules": schedules},
headers=_auth_header(),
)
assert resp.status_code == 200
data = resp.json()
assert len(data) == 2
assert data[0]["start_time"] == "08:00"
assert data[0]["provider_fee"] == 1.05
@pytest.mark.integration
@pytest.mark.asyncio
async def test_put_fee_schedules_persisted(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
"""Saved schedules are returned by a subsequent GET."""
provider_id = await _create_provider(integration_client)
schedules = [{"start_time": "09:00", "end_time": "17:00", "provider_fee": 1.07}]
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={"schedules": schedules},
headers=_auth_header(),
)
get_resp = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
headers=_auth_header(),
)
assert get_resp.status_code == 200
assert get_resp.json()[0]["provider_fee"] == 1.07
@pytest.mark.integration
@pytest.mark.asyncio
async def test_put_fee_schedules_replaces_existing(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
# Set initial schedule
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "08:00", "end_time": "12:00", "provider_fee": 1.03}
]
},
headers=_auth_header(),
)
# Replace with different schedule
resp = await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "14:00", "end_time": "20:00", "provider_fee": 1.08}
]
},
headers=_auth_header(),
)
assert resp.status_code == 200
data = resp.json()
assert len(data) == 1
assert data[0]["start_time"] == "14:00"
@pytest.mark.integration
@pytest.mark.asyncio
async def test_put_fee_schedules_overlap_rejected(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
schedules = [
{"start_time": "08:00", "end_time": "14:00", "provider_fee": 1.05},
{"start_time": "12:00", "end_time": "18:00", "provider_fee": 1.03},
]
resp = await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={"schedules": schedules},
headers=_auth_header(),
)
assert resp.status_code == 400
assert "overlap" in resp.json()["detail"].lower()
@pytest.mark.integration
@pytest.mark.asyncio
async def test_put_fee_schedules_invalid_time_format_rejected(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
resp = await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "8:00", "end_time": "18:00", "provider_fee": 1.05}
]
},
headers=_auth_header(),
)
assert resp.status_code == 400
@pytest.mark.integration
@pytest.mark.asyncio
async def test_put_fee_schedules_invalid_fee_rejected(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
resp = await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "08:00", "end_time": "18:00", "provider_fee": -0.5}
]
},
headers=_auth_header(),
)
assert resp.status_code == 400
@pytest.mark.integration
@pytest.mark.asyncio
async def test_put_fee_schedules_empty_clears_schedules(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
# Set a schedule
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05}
]
},
headers=_auth_header(),
)
# Clear with empty list
resp = await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={"schedules": []},
headers=_auth_header(),
)
assert resp.status_code == 200
assert resp.json() == []
@pytest.mark.integration
@pytest.mark.asyncio
async def test_put_fee_schedules_404_for_missing_provider(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
resp = await integration_client.put(
"/admin/api/upstream-providers/99999/fee-schedules",
json={"schedules": []},
headers=_auth_header(),
)
assert resp.status_code == 404
# ---------------------------------------------------------------------------
# DELETE fee schedules
# ---------------------------------------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_delete_fee_schedules(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
# Add schedules
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05}
]
},
headers=_auth_header(),
)
# Delete
del_resp = await integration_client.delete(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
headers=_auth_header(),
)
assert del_resp.status_code == 200
assert del_resp.json()["ok"] is True
# Verify schedules are gone
get_resp = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
headers=_auth_header(),
)
assert get_resp.json() == []
@pytest.mark.integration
@pytest.mark.asyncio
async def test_delete_fee_schedules_404_for_missing_provider(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
resp = await integration_client.delete(
"/admin/api/upstream-providers/99999/fee-schedules",
headers=_auth_header(),
)
assert resp.status_code == 404
# ---------------------------------------------------------------------------
# Fee schedules appear in provider list and detail
# ---------------------------------------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_fee_schedules_in_provider_list(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05}
]
},
headers=_auth_header(),
)
list_resp = await integration_client.get(
"/admin/api/upstream-providers", headers=_auth_header()
)
assert list_resp.status_code == 200
providers = list_resp.json()
target = next((p for p in providers if p["id"] == provider_id), None)
assert target is not None
assert len(target["provider_fee_schedules"]) == 1
assert target["provider_fee_schedules"][0]["provider_fee"] == 1.05
@pytest.mark.integration
@pytest.mark.asyncio
async def test_fee_schedules_in_provider_detail(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "10:00", "end_time": "22:00", "provider_fee": 1.06}
]
},
headers=_auth_header(),
)
detail_resp = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}", headers=_auth_header()
)
assert detail_resp.status_code == 200
data = detail_resp.json()
assert len(data["provider_fee_schedules"]) == 1
assert data["provider_fee_schedules"][0]["start_time"] == "10:00"
# ---------------------------------------------------------------------------
# Provider deletion clears fee schedules
# ---------------------------------------------------------------------------
@pytest.mark.integration
@pytest.mark.asyncio
async def test_provider_delete_clears_fee_schedules(
integration_client: AsyncClient, patched_db_engine: Any
) -> None:
provider_id = await _create_provider(integration_client)
await integration_client.put(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
json={
"schedules": [
{"start_time": "08:00", "end_time": "18:00", "provider_fee": 1.05}
]
},
headers=_auth_header(),
)
# Delete provider
del_resp = await integration_client.delete(
f"/admin/api/upstream-providers/{provider_id}", headers=_auth_header()
)
assert del_resp.status_code == 200
# Provider is gone → schedule endpoint returns 404
get_resp = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/fee-schedules",
headers=_auth_header(),
)
assert get_resp.status_code == 404

View File

@@ -0,0 +1,268 @@
"""
Tests for the reservation lifecycle:
1. Reserve → reserved_balance increases, available (total_balance) decreases.
2. Reserve → revert → reserved_balance restored, balance untouched.
3. Reserve → finalise → reserved_balance released, balance charged.
4. Two parallel reserves, only one fits → second blocked with 402.
5. Three parallel reserves, two fit, third blocked with 402.
6. Sequential reserves until balance exhausted → next request blocked.
Reservation invariant enforced by the atomic WHERE clause in pay_for_request:
balance - reserved_balance >= cost_per_request
"""
import asyncio
import uuid
import pytest
from fastapi import HTTPException
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.auth import pay_for_request, revert_pay_for_request
from routstr.core.db import ApiKey, create_session
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_key(balance: int, reserved: int = 0) -> ApiKey:
return ApiKey(
hashed_key=f"test_{uuid.uuid4().hex}",
balance=balance,
reserved_balance=reserved,
total_spent=0,
total_requests=0,
)
async def _persist(session: AsyncSession, key: ApiKey) -> ApiKey:
session.add(key)
await session.commit()
await session.refresh(key)
return key
# ---------------------------------------------------------------------------
# Test 1 — Reserve: reserved_balance increases, available balance decreases
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_reserve_increases_reserved_balance(
integration_session: AsyncSession,
) -> None:
"""pay_for_request must increment reserved_balance by cost_per_request."""
cost = 100
key = await _persist(integration_session, _make_key(balance=500))
await pay_for_request(key, cost, integration_session)
await integration_session.refresh(key)
assert key.reserved_balance == cost
assert key.balance == 500 # balance column is NOT decremented on reserve
assert key.total_balance == 500 - cost # available = balance - reserved
# ---------------------------------------------------------------------------
# Test 2 — Revert: reserved_balance restored, balance untouched
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_revert_releases_reservation(
integration_session: AsyncSession,
) -> None:
"""revert_pay_for_request must release the reservation without touching balance."""
cost = 150
key = await _persist(integration_session, _make_key(balance=300))
await pay_for_request(key, cost, integration_session)
await integration_session.refresh(key)
assert key.reserved_balance == cost
await revert_pay_for_request(key, integration_session, cost)
await integration_session.refresh(key)
assert key.reserved_balance == 0
assert key.balance == 300 # balance unchanged after revert
# ---------------------------------------------------------------------------
# Test 3 — Finalise: reservation released + balance charged
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_finalise_releases_reservation_and_charges_balance(
integration_session: AsyncSession,
) -> None:
"""adjust_payment_for_tokens must zero reserved_balance and deduct actual cost."""
from unittest.mock import patch
from routstr.auth import adjust_payment_for_tokens
from routstr.payment.cost_calculation import CostData
cost = 100
actual = 80 # actual < reserved → refund path
key = await _persist(integration_session, _make_key(balance=500))
await pay_for_request(key, cost, integration_session)
await integration_session.refresh(key)
assert key.reserved_balance == cost
cost_data = CostData(
base_msats=0,
input_msats=40,
output_msats=40,
total_msats=actual,
total_usd=0.0,
input_tokens=50,
output_tokens=50,
)
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
with patch("routstr.auth.calculate_cost", return_value=cost_data):
await adjust_payment_for_tokens(key, response_data, integration_session, cost)
await integration_session.refresh(key)
assert key.reserved_balance == 0
assert key.balance == 500 - actual
assert key.total_spent == actual
# ---------------------------------------------------------------------------
# Test 4 — Concurrent: second parallel reserve blocked when balance exhausted
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_concurrent_second_reserve_blocked_when_balance_exhausted(
patched_db_engine: None,
) -> None:
"""When two requests race for the same balance, only one succeeds; the other gets 402."""
cost = 300
key_hash = f"test_concurrent_{uuid.uuid4().hex}"
async with create_session() as session:
key = ApiKey(
hashed_key=key_hash,
balance=300, # exactly enough for ONE reservation
reserved_balance=0,
total_spent=0,
total_requests=0,
)
session.add(key)
await session.commit()
results: list[str] = []
async def attempt_reserve() -> None:
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)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(attempt_reserve(), attempt_reserve())
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, key_hash)
assert final is not None
# reserved_balance must equal exactly one reservation (not two)
assert final.reserved_balance == cost, (
f"Expected reserved_balance={cost}, got {final.reserved_balance}"
)
assert final.balance == 300, "Balance column must not be modified by reservation"
# ---------------------------------------------------------------------------
# Test 5 — Concurrent: three requests, two fit, third blocked
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_three_parallel_reserves_third_blocked(
patched_db_engine: None,
) -> None:
"""Balance covers two reservations exactly; the third concurrent request must be blocked."""
cost = 100
key_hash = f"test_three_parallel_{uuid.uuid4().hex}"
async with create_session() as session:
key = ApiKey(
hashed_key=key_hash,
balance=200, # fits exactly 2 reservations of 100
reserved_balance=0,
total_spent=0,
total_requests=0,
)
session.add(key)
await session.commit()
results: list[str] = []
async def attempt_reserve() -> None:
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)
results.append("success")
except HTTPException as exc:
assert exc.status_code == 402
results.append("blocked")
await asyncio.gather(
attempt_reserve(),
attempt_reserve(),
attempt_reserve(),
)
successes = results.count("success")
blocked = results.count("blocked")
assert successes == 2, f"Expected 2 successes, got {successes}: {results}"
assert blocked == 1, f"Expected 1 blocked, got {blocked}: {results}"
async with create_session() as session:
final = await session.get(ApiKey, key_hash)
assert final is not None
assert final.reserved_balance == cost * 2, (
f"Expected reserved_balance={cost * 2}, got {final.reserved_balance}"
)
assert final.balance == 200, "Balance column must not be modified by reservation"
# ---------------------------------------------------------------------------
# Test 6 — Sequential exhaustion: reserve until empty, next request blocked
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_sequential_reserves_block_when_balance_exhausted(
integration_session: AsyncSession,
) -> None:
"""Repeated reservations should block as soon as available balance drops below cost."""
cost = 100
key = await _persist(integration_session, _make_key(balance=250))
# First two succeed (100 + 100 = 200 ≤ 250)
await pay_for_request(key, cost, integration_session)
await pay_for_request(key, cost, integration_session)
await integration_session.refresh(key)
assert key.reserved_balance == 200
assert key.total_balance == 50 # 250 - 200
# Third: only 50 available, need 100 → blocked
with pytest.raises(HTTPException) as exc_info:
await pay_for_request(key, cost, integration_session)
assert exc_info.value.status_code == 402
await integration_session.refresh(key)
assert key.reserved_balance == 200 # unchanged after failed reserve

View File

@@ -13,7 +13,7 @@ import pytest
from httpx import AsyncClient
from sqlmodel import select
from routstr.core.db import ApiKey
from routstr.core.db import ApiKey, CashuTransaction
@pytest.mark.integration
@@ -356,6 +356,70 @@ async def test_concurrent_refund_requests(
assert len(successful) + len(failed) == 5
@pytest.mark.integration
@pytest.mark.asyncio
async def test_refund_rejects_concurrent_topup_on_same_key(
authenticated_client: AsyncClient,
testmint_wallet: Any,
) -> None:
"""Test refund returns 409 when a concurrent topup changes the balance first."""
from routstr import balance as balance_module
wallet_response = await authenticated_client.get("/v1/wallet/")
assert wallet_response.status_code == 200
initial_balance = wallet_response.json()["balance"]
topup_amount_sat = 500
topup_token = await testmint_wallet.mint_tokens(topup_amount_sat)
validate_called = asyncio.Event()
allow_refund_to_continue = asyncio.Event()
original_validate_bearer_key = balance_module.validate_bearer_key
delayed_once = False
async def delayed_validate_bearer_key(*args: Any, **kwargs: Any) -> ApiKey:
nonlocal delayed_once
key = await original_validate_bearer_key(*args, **kwargs)
if not delayed_once:
delayed_once = True
validate_called.set()
await allow_refund_to_continue.wait()
return key
async def issue_refund() -> Any:
return await authenticated_client.post("/v1/wallet/refund")
async def issue_topup() -> Any:
await validate_called.wait()
try:
return await authenticated_client.post(
"/v1/wallet/topup", params={"cashu_token": topup_token}
)
finally:
allow_refund_to_continue.set()
with patch(
"routstr.balance.validate_bearer_key", new=delayed_validate_bearer_key
):
refund_response, topup_response = await asyncio.gather(
issue_refund(), issue_topup()
)
assert topup_response.status_code == 200
assert topup_response.json()["msats"] == topup_amount_sat * 1000
assert refund_response.status_code == 409
assert (
refund_response.json()["detail"]
== "Balance changed concurrently. Please retry the refund."
)
final_balance_response = await authenticated_client.get("/v1/wallet/")
assert final_balance_response.status_code == 200
assert final_balance_response.json()["balance"] == (
initial_balance + topup_amount_sat * 1000
)
@pytest.mark.integration
@pytest.mark.asyncio
async def test_refund_during_active_usage(
@@ -394,6 +458,42 @@ async def test_refund_during_active_usage(
assert response.json()["balance"] == 0
@pytest.mark.integration
@pytest.mark.asyncio
async def test_wallet_history_returns_apikey_transactions(
authenticated_client: AsyncClient,
testmint_wallet: Any,
integration_session: Any,
) -> None:
wallet_response = await authenticated_client.get("/v1/wallet/")
api_key = wallet_response.json()["api_key"]
hashed_key = api_key[3:] if api_key.startswith("sk-") else api_key
topup_token = await testmint_wallet.mint_tokens(250)
topup_response = await authenticated_client.post(
"/v1/wallet/topup", params={"cashu_token": topup_token}
)
assert topup_response.status_code == 200
refund_response = await authenticated_client.post("/v1/wallet/refund")
assert refund_response.status_code == 200
history_response = await authenticated_client.get("/v1/wallet/history")
assert history_response.status_code == 200
transactions = history_response.json()["transactions"]
assert len(transactions) >= 2
assert all("api_key_hashed_key" not in tx for tx in transactions)
assert {tx["type"] for tx in transactions} >= {"in", "out"}
db_result = await integration_session.execute(
select(CashuTransaction).where(
CashuTransaction.api_key_hashed_key == hashed_key
)
)
db_transactions = db_result.scalars().all()
assert len(db_transactions) >= 2
@pytest.mark.integration
@pytest.mark.asyncio
async def test_mint_unavailability_handling(

View File

@@ -11,7 +11,7 @@ import pytest
from httpx import AsyncClient
from sqlmodel import select
from routstr.core.db import ApiKey
from routstr.core.db import ApiKey, CashuTransaction
from .utils import (
CashuTokenGenerator,
@@ -71,6 +71,16 @@ async def test_topup_with_valid_token( # type: ignore[no-untyped-def]
assert db_key.balance == new_balance
assert db_key.balance == initial_balance + (topup_amount * 1000)
tx_result = await integration_session.execute(
select(CashuTransaction).where(
CashuTransaction.token == token,
CashuTransaction.type == "in",
)
)
tx = tx_result.scalar_one()
assert tx.api_key_hashed_key == hashed_key
assert tx.source == "apikey"
@pytest.mark.integration
@pytest.mark.asyncio

View File

@@ -1,11 +1,12 @@
import json
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi.responses import JSONResponse
from routstr.balance import refund_wallet_endpoint
from routstr.core.db import CashuTransaction
from routstr.core.db import ApiKey, CashuTransaction
from routstr.wallet import credit_balance
def _make_cashu_tx(
@@ -29,6 +30,12 @@ def _exec_result(tx: CashuTransaction | None) -> MagicMock:
return result
def _update_result(rowcount: int) -> MagicMock:
result = MagicMock()
result.rowcount = rowcount
return result
@pytest.mark.asyncio
async def test_refund_x_cashu_returns_token() -> None:
x_cashu_token = "cashuAtest_token_value"
@@ -114,3 +121,237 @@ async def test_refund_x_cashu_swept_raises_410() -> None:
)
assert exc_info.value.status_code == 410
# ---------------------------------------------------------------------------
# source field defaults
# ---------------------------------------------------------------------------
def test_cashu_transaction_source_defaults_to_x_cashu() -> None:
tx = CashuTransaction(token="cashuAtest", amount=100, unit="msat")
assert tx.source == "x-cashu"
def test_cashu_transaction_source_can_be_apikey() -> None:
tx = CashuTransaction(token="cashuAtest", amount=100, unit="msat", source="apikey")
assert tx.source == "apikey"
# ---------------------------------------------------------------------------
# apikey-based refund: token logging and CashuTransaction storage
# ---------------------------------------------------------------------------
def _make_api_key(
balance: int = 5000,
refund_currency: str | None = "sat",
refund_mint_url: str | None = "https://mint.example.com",
refund_address: str | None = None,
parent_key_hash: str | None = None,
) -> ApiKey:
key = ApiKey(hashed_key="testhash")
key.balance = balance
key.reserved_balance = 0
key.refund_currency = refund_currency
key.refund_mint_url = refund_mint_url
key.refund_address = refund_address
key.parent_key_hash = parent_key_hash
key.total_spent = 0
key.total_requests = 0
return key
@pytest.mark.asyncio
async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> None:
key = _make_api_key(balance=5000, refund_currency="sat")
refund_token = "cashuArefund_apikey_token"
session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock()
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store,
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
):
result = await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
assert isinstance(result, dict)
assert result["token"] == refund_token
mock_store.assert_awaited_once()
call_kwargs = mock_store.call_args.kwargs
assert call_kwargs["source"] == "apikey"
assert call_kwargs["token"] == refund_token
assert call_kwargs["typ"] == "out"
assert call_kwargs["api_key_hashed_key"] == key.hashed_key
@pytest.mark.asyncio
async def test_apikey_refund_logs_token() -> None:
key = _make_api_key(balance=5000, refund_currency="sat")
refund_token = "cashuAlogged_token"
session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock()
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger") as mock_logger,
):
await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
calls = [str(c) for c in mock_logger.info.call_args_list]
assert any("cashu token issued" in c for c in calls)
@pytest.mark.asyncio
async def test_apikey_refund_log_includes_path() -> None:
key = _make_api_key(balance=5000, refund_currency="sat")
refund_token = "cashuApath_token"
session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock()
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger") as mock_logger,
):
await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
# Find the "cashu token issued" call and verify extra contains the path
token_issued_calls = [
c for c in mock_logger.info.call_args_list
if c.args and "cashu token issued" in c.args[0]
]
assert len(token_issued_calls) == 1
extra = token_issued_calls[0].kwargs.get("extra", {})
assert extra.get("path") == "/v1/wallet/refund"
@pytest.mark.asyncio
async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None:
"""When the debit CAS fails (rowcount=0), no token is minted and 409 is returned."""
from fastapi import HTTPException
key = _make_api_key(balance=5000, refund_currency="sat")
session = MagicMock()
# Debit returns rowcount=0 → balance changed concurrently
session.exec = AsyncMock(return_value=_update_result(0))
session.commit = AsyncMock()
mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted")
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", mock_send_token),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 409
# Crucially: send_token must NOT have been called
mock_send_token.assert_not_awaited()
@pytest.mark.asyncio
async def test_credit_balance_stores_apikey_transaction_history() -> None:
key = _make_api_key(balance=1000)
session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1))
session.commit = AsyncMock()
session.refresh = AsyncMock()
with (
patch(
"routstr.wallet.recieve_token",
AsyncMock(return_value=(100, "sat", "https://mint.example")),
),
patch("routstr.wallet.store_cashu_transaction", AsyncMock()) as mock_store,
):
amount = await credit_balance("cashuAtopup_token", key, session)
assert amount == 100_000
mock_store.assert_awaited_once()
call_kwargs = mock_store.call_args.kwargs
assert call_kwargs["typ"] == "in"
assert call_kwargs["source"] == "apikey"
assert call_kwargs["api_key_hashed_key"] == key.hashed_key
assert call_kwargs["amount"] == 100
assert call_kwargs["unit"] == "sat"
assert call_kwargs["token"] == "cashuAtopup_token"
assert call_kwargs["mint_url"] == "https://mint.example"
@pytest.mark.asyncio
async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
"""When debit succeeds but minting fails, balance must be restored."""
from fastapi import HTTPException
key = _make_api_key(balance=5000, refund_currency="sat")
# First exec call = debit (succeeds), second = restore
session = MagicMock()
session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)])
session.commit = AsyncMock()
with (
patch("routstr.balance.validate_bearer_key", AsyncMock(return_value=key)),
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(side_effect=Exception("mint down"))),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger"),
):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 503
# Verify two exec calls: debit + restore
assert session.exec.await_count == 2

View File

@@ -0,0 +1,256 @@
"""Unit tests for routstr.payment.fee_schedule."""
from datetime import datetime, timedelta, timezone
import pytest
from pydantic.v1 import ValidationError
from routstr.payment.fee_schedule import (
FeeTimeRange,
get_active_fee,
ranges_overlap,
validate_no_overlaps,
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _r(start: str, end: str, fee: float = 1.05) -> FeeTimeRange:
return FeeTimeRange(start_time=start, end_time=end, provider_fee=fee)
def _now(h: int, m: int = 0) -> datetime:
return datetime(2026, 1, 1, h, m, tzinfo=timezone.utc)
# ---------------------------------------------------------------------------
# FeeTimeRange validation
# ---------------------------------------------------------------------------
class TestFeeTimeRangeValidation:
def test_valid_range(self) -> None:
r = _r("08:00", "18:00", 1.05)
assert r.start_time == "08:00"
assert r.end_time == "18:00"
assert r.provider_fee == 1.05
def test_invalid_start_time_format(self) -> None:
with pytest.raises(ValidationError, match="HH:MM"):
_r("8:00", "18:00")
def test_invalid_end_time_hour_out_of_range(self) -> None:
with pytest.raises(ValidationError):
_r("08:00", "24:00")
def test_invalid_end_time_minute_out_of_range(self) -> None:
with pytest.raises(ValidationError):
_r("08:00", "18:60")
def test_invalid_time_letters(self) -> None:
with pytest.raises(ValidationError):
_r("ab:cd", "18:00")
def test_fee_must_be_positive(self) -> None:
with pytest.raises(ValidationError, match="provider_fee must be > 0"):
_r("08:00", "18:00", fee=0.0)
def test_fee_negative_rejected(self) -> None:
with pytest.raises(ValidationError):
_r("08:00", "18:00", fee=-0.5)
def test_fee_below_one_allowed(self) -> None:
r = _r("08:00", "18:00", fee=0.95)
assert r.provider_fee == 0.95
def test_boundary_times_valid(self) -> None:
r = _r("00:00", "23:59")
assert r.start_time == "00:00"
assert r.end_time == "23:59"
# ---------------------------------------------------------------------------
# ranges_overlap
# ---------------------------------------------------------------------------
class TestRangesOverlap:
def test_non_overlapping_ranges(self) -> None:
assert not ranges_overlap(_r("08:00", "12:00"), _r("12:00", "18:00"))
def test_overlapping_ranges(self) -> None:
assert ranges_overlap(_r("08:00", "14:00"), _r("12:00", "18:00"))
def test_one_contains_the_other(self) -> None:
assert ranges_overlap(_r("08:00", "20:00"), _r("10:00", "18:00"))
def test_identical_ranges_overlap(self) -> None:
assert ranges_overlap(_r("08:00", "12:00"), _r("08:00", "12:00"))
def test_adjacent_non_overlapping(self) -> None:
# end of first == start of second → no overlap (open interval [start, end))
assert not ranges_overlap(_r("06:00", "12:00"), _r("12:00", "18:00"))
def test_midnight_crossing_vs_day_range_overlap(self) -> None:
# 22:0006:00 crosses midnight; 04:0008:00 should overlap (both cover 04:0006:00)
assert ranges_overlap(_r("22:00", "06:00"), _r("04:00", "08:00"))
def test_midnight_crossing_vs_non_overlapping_day_range(self) -> None:
# 22:0006:00 does NOT cover 10:0018:00
assert not ranges_overlap(_r("22:00", "06:00"), _r("10:00", "18:00"))
def test_two_midnight_crossing_ranges_overlap(self) -> None:
assert ranges_overlap(_r("20:00", "04:00"), _r("22:00", "06:00"))
def test_two_midnight_crossing_ranges_non_overlap(self) -> None:
# 21:0023:00 and 23:0021:00 (full day minus one hour): they do overlap
# Let's use a case that genuinely doesn't: 21:0022:00 adjacent
# Actually for two midnight-crossing ranges it's hard to not overlap—let's test equal endpoints
assert not ranges_overlap(_r("22:00", "23:00"), _r("23:00", "01:00"))
# ---------------------------------------------------------------------------
# validate_no_overlaps
# ---------------------------------------------------------------------------
class TestValidateNoOverlaps:
def test_no_overlaps_passes(self) -> None:
validate_no_overlaps(
[_r("00:00", "08:00"), _r("08:00", "16:00"), _r("16:00", "23:59")]
)
def test_overlap_raises(self) -> None:
with pytest.raises(ValueError, match="overlap"):
validate_no_overlaps([_r("08:00", "14:00"), _r("12:00", "18:00")])
def test_single_range_passes(self) -> None:
validate_no_overlaps([_r("08:00", "18:00")])
def test_empty_list_passes(self) -> None:
validate_no_overlaps([])
def test_midnight_crossing_overlap_detected(self) -> None:
with pytest.raises(ValueError, match="overlap"):
validate_no_overlaps([_r("22:00", "06:00"), _r("04:00", "08:00")])
# ---------------------------------------------------------------------------
# get_active_fee
# ---------------------------------------------------------------------------
class TestGetActiveFee:
def test_returns_default_when_no_ranges(self) -> None:
assert get_active_fee(None, 1.01) == 1.01
def test_returns_default_for_empty_list(self) -> None:
assert get_active_fee([], 1.01) == 1.01
def test_returns_matching_fee(self) -> None:
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=_now(12)) == 1.05
def test_returns_default_when_no_match(self) -> None:
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=_now(20)) == 1.01
def test_boundary_start_inclusive(self) -> None:
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=_now(8, 0)) == 1.05
def test_boundary_end_exclusive(self) -> None:
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=_now(18, 0)) == 1.01
def test_midnight_crossing_before_midnight(self) -> None:
ranges = [_r("22:00", "06:00", fee=1.03)]
assert get_active_fee(ranges, 1.01, _now=_now(23)) == 1.03
def test_midnight_crossing_after_midnight(self) -> None:
ranges = [_r("22:00", "06:00", fee=1.03)]
assert get_active_fee(ranges, 1.01, _now=_now(3)) == 1.03
def test_midnight_crossing_outside_range(self) -> None:
ranges = [_r("22:00", "06:00", fee=1.03)]
assert get_active_fee(ranges, 1.01, _now=_now(12)) == 1.01
def test_multiple_ranges_correct_match(self) -> None:
ranges = [
_r("00:00", "08:00", fee=1.02),
_r("08:00", "16:00", fee=1.05),
_r("16:00", "23:59", fee=1.03),
]
assert get_active_fee(ranges, 1.01, _now=_now(10)) == 1.05
assert get_active_fee(ranges, 1.01, _now=_now(2)) == 1.02
assert get_active_fee(ranges, 1.01, _now=_now(20)) == 1.03
def test_first_matching_range_wins(self) -> None:
# When multiple ranges could match (should not happen if validated),
# the first one wins.
ranges = [_r("08:00", "20:00", fee=1.05), _r("10:00", "12:00", fee=1.02)]
assert get_active_fee(ranges, 1.01, _now=_now(11)) == 1.05
# ---------------------------------------------------------------------------
# Timezone-aware inputs (CEST / CET)
# ---------------------------------------------------------------------------
class TestGetActiveFeeTimezones:
"""Verify that tz-aware datetimes are normalised to UTC before matching."""
# CEST = UTC+2 (Central European Summer Time, used ~late March late Oct)
CEST = timezone(timedelta(hours=2))
# CET = UTC+1 (Central European Time, used the rest of the year)
CET = timezone(timedelta(hours=1))
def test_cest_datetime_normalised_to_utc_matches(self) -> None:
# 10:00 CEST == 08:00 UTC — schedule 08:0018:00 should match
now_cest = datetime(2026, 7, 1, 10, 0, tzinfo=self.CEST)
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.05
def test_cest_datetime_normalised_to_utc_no_match(self) -> None:
# 06:00 CEST == 04:00 UTC — schedule 08:0018:00 should NOT match
now_cest = datetime(2026, 7, 1, 6, 0, tzinfo=self.CEST)
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.01
def test_cet_datetime_normalised_to_utc_matches(self) -> None:
# 09:00 CET == 08:00 UTC — schedule 08:0018:00 should match
now_cet = datetime(2026, 1, 15, 9, 0, tzinfo=self.CET)
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=now_cet) == 1.05
def test_cet_datetime_before_utc_range(self) -> None:
# 08:30 CET == 07:30 UTC — schedule 08:0018:00 should NOT match
now_cet = datetime(2026, 1, 15, 8, 30, tzinfo=self.CET)
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=now_cet) == 1.01
def test_cest_midnight_crossing_before_midnight(self) -> None:
# 00:30 CEST == 22:30 UTC — schedule 22:0006:00 UTC should match
now_cest = datetime(2026, 7, 2, 0, 30, tzinfo=self.CEST)
ranges = [_r("22:00", "06:00", fee=1.03)]
assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.03
def test_cest_midnight_crossing_after_midnight(self) -> None:
# 05:00 CEST == 03:00 UTC — schedule 22:0006:00 UTC should match
now_cest = datetime(2026, 7, 2, 5, 0, tzinfo=self.CEST)
ranges = [_r("22:00", "06:00", fee=1.03)]
assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.03
def test_cest_midnight_crossing_outside_range(self) -> None:
# 14:00 CEST == 12:00 UTC — schedule 22:0006:00 UTC should NOT match
now_cest = datetime(2026, 7, 2, 14, 0, tzinfo=self.CEST)
ranges = [_r("22:00", "06:00", fee=1.03)]
assert get_active_fee(ranges, 1.01, _now=now_cest) == 1.01
def test_naive_utc_datetime_still_works(self) -> None:
# Naive datetimes are treated as UTC (defensive fallback path)
now_naive = datetime(2026, 1, 1, 12, 0) # no tzinfo
ranges = [_r("08:00", "18:00", fee=1.05)]
assert get_active_fee(ranges, 1.01, _now=now_naive) == 1.05

View File

@@ -0,0 +1,117 @@
"""Regression tests for the periodic upstream models refresh loop."""
from __future__ import annotations
import asyncio
import os
from typing import cast
from unittest.mock import AsyncMock
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
class _FakeUpstream:
"""Minimal stand-in for BaseUpstreamProvider used by the refresh loop.
Only ``base_url`` (for error logging) and ``refresh_models_cache`` (the call
under test) are exercised; everything else stays unused.
"""
def __init__(self, name: str) -> None:
self.base_url = f"http://{name}"
self.refresh_models_cache = AsyncMock()
def _make_fake_upstream(name: str) -> BaseUpstreamProvider:
# The loop only uses duck-typed attributes — cast keeps the test type-clean
# without dragging in BaseUpstreamProvider's full constructor.
return cast(BaseUpstreamProvider, _FakeUpstream(name))
@pytest.mark.asyncio
async def test_refresh_loop_picks_up_providers_added_after_startup(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""If a provider is added after the loop starts (e.g. via reinitialize_upstreams),
the next loop iteration must refresh it. Previously the loop captured the upstream
list at startup and missed any later additions."""
from routstr.core.settings import settings as global_settings
from routstr.upstream.helpers import refresh_upstreams_models_periodically
# Tight interval so the test finishes quickly.
monkeypatch.setattr(
global_settings, "models_refresh_interval_seconds", 1, raising=False
)
initial_upstream = _make_fake_upstream("initial")
live_list: list[BaseUpstreamProvider] = [initial_upstream]
# Stub out the post-iteration sats-pricing refresh so the loop body has no DB deps.
async def _noop_pricing_refresh() -> None: # pragma: no cover - trivial stub
return None
monkeypatch.setattr(
"routstr.payment.models._update_sats_pricing_once",
_noop_pricing_refresh,
)
task = asyncio.create_task(
refresh_upstreams_models_periodically(lambda: live_list)
)
try:
# Wait for the first iteration to refresh the initial upstream.
for _ in range(40):
if initial_upstream.refresh_models_cache.await_count >= 1: # type: ignore[attr-defined]
break
await asyncio.sleep(0.05)
assert initial_upstream.refresh_models_cache.await_count >= 1, ( # type: ignore[attr-defined]
"loop did not refresh the initial upstream within the timeout"
)
# Simulate reinitialize_upstreams: replace the live list contents with new
# provider instances. The loop must observe the swap on its next tick.
new_upstream = _make_fake_upstream("added-after-startup")
live_list[:] = [new_upstream]
for _ in range(60):
if new_upstream.refresh_models_cache.await_count >= 1: # type: ignore[attr-defined]
break
await asyncio.sleep(0.05)
assert new_upstream.refresh_models_cache.await_count >= 1, ( # type: ignore[attr-defined]
"loop did not refresh the upstream added after startup — "
"regression: list snapshot captured at startup"
)
finally:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
@pytest.mark.asyncio
async def test_refresh_loop_disabled_when_interval_non_positive(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from routstr.core.settings import settings as global_settings
from routstr.upstream.helpers import refresh_upstreams_models_periodically
monkeypatch.setattr(
global_settings, "models_refresh_interval_seconds", 0, raising=False
)
upstream = _make_fake_upstream("never-refreshed")
# Loop must return immediately without ever touching the upstream.
await asyncio.wait_for(
refresh_upstreams_models_periodically(lambda: [upstream]),
timeout=1.0,
)
upstream.refresh_models_cache.assert_not_awaited() # type: ignore[attr-defined]

View File

@@ -0,0 +1,117 @@
import json
from collections.abc import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock
import pytest
from routstr.core.db import ApiKey
from routstr.upstream.base import BaseUpstreamProvider
@pytest.mark.asyncio
async def test_stream_with_id_injection() -> None:
"""Test that stream_with_cost correctly injects IDs into complete JSON chunks but skips partials."""
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test_key"
)
# Mock response with mixed chunks:
# 1. Complete JSON without ID
# 2. Partial JSON (should be passed through)
# 3. Complete JSON with ID (should be preserved or updated if requested_model is set)
# 4. [DONE] message
chunks = [
b'data: {"choices": [{"delta": {"content": "Hello"}}]}\n\n',
b'data: {"choices": [{"delta": {"content": "', # Partial
b'world"}}]}\n\n',
b'data: {"id": "existing-id", "choices": [{"delta": {"content": "!"}}]}\n\n',
b"data: [DONE]\n\n",
]
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
for chunk in chunks:
yield chunk
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"}
mock_response.aiter_bytes = aiter_bytes
key = MagicMock(spec=ApiKey)
key.hashed_key = "test_hash"
key.balance = 1000
background_tasks = MagicMock()
# We need to mock adjust_payment_for_tokens since it's called at the end
with MagicMock():
from routstr.upstream import base
# Mocking the module-level function used in the generator
base.adjust_payment_for_tokens = AsyncMock(
return_value={"total_usd": 0.1, "total_msats": 100}
)
# create_session() is used as an async context manager whose entered
# value exposes an awaitable .get(). Build a mock that behaves that
# way so the post-stream cost-chunk emission can run.
mock_session = MagicMock()
mock_session.get = AsyncMock(return_value=key)
mock_ctx = MagicMock()
mock_ctx.__aenter__ = AsyncMock(return_value=mock_session)
mock_ctx.__aexit__ = AsyncMock(return_value=None)
base.create_session = MagicMock(return_value=mock_ctx)
streaming_response = await provider.handle_streaming_chat_completion(
response=mock_response,
key=key,
max_cost_for_model=100,
background_tasks=background_tasks,
requested_model="test-model",
)
results = []
async for chunk in streaming_response.body_iterator:
results.append(chunk)
# Parse results
parsed_results = []
for r in results:
if isinstance(r, bytes) and r.startswith(b"data: "):
data = r[6:].decode().strip()
if data == "[DONE]":
parsed_results.append(data)
else:
try:
parsed_results.append(json.loads(data))
except (json.JSONDecodeError, UnicodeDecodeError):
parsed_results.append(
data
) # Keep as string if it failed to parse
# Verifications
# 1. First chunk should have an injected ID and the requested model
assert isinstance(parsed_results[0], dict)
assert "id" in parsed_results[0]
assert parsed_results[0]["id"].startswith("chatcmpl-")
assert parsed_results[0]["model"] == "test-model"
# 2. Second chunk was partial, should be passed as-is
# In current implementation, re.split(b"data: ", b'data: {...') gives ['', '{...']
# The first empty part is skipped. The second part is processed.
# Check that we have results
assert len(parsed_results) >= 4
# Find the chunk that was "existing-id"
id_chunk = next(
r
for r in parsed_results
if isinstance(r, dict)
and "choices" in r
and r["choices"][0]["delta"].get("content") == "!"
)
assert id_chunk["id"] == parsed_results[0]["id"]
assert id_chunk["model"] == "test-model"
# 4. [DONE] should be there
assert "[DONE]" in parsed_results

View File

@@ -220,6 +220,33 @@ async def test_recieve_token_untrusted_mint() -> None:
@pytest.mark.asyncio
@pytest.mark.asyncio
async def test_swap_to_primary_mint_already_on_primary() -> None:
from routstr.core.settings import settings
from routstr.wallet import swap_to_primary_mint
mock_token = Mock()
mock_token.mint = settings.primary_mint
mock_token.amount = 1000
mock_token.unit = "sat"
mock_token.proofs = []
mock_token_wallet = Mock()
mock_token_wallet.split = AsyncMock(return_value=None)
mock_token_wallet.request_mint = AsyncMock()
mock_token_wallet.melt_quote = AsyncMock()
with patch("routstr.wallet.get_wallet", AsyncMock(return_value=mock_token_wallet)):
amount, unit, mint = await swap_to_primary_mint(mock_token, mock_token_wallet)
assert amount == 1000
assert unit == "sat"
assert mint == settings.primary_mint
mock_token_wallet.split.assert_called_once()
mock_token_wallet.request_mint.assert_not_called()
mock_token_wallet.melt_quote.assert_not_called()
async def test_swap_to_primary_mint_success() -> None:
"""Test successful swap with dynamic fee calculation."""
from routstr.wallet import swap_to_primary_mint

View File

@@ -0,0 +1,216 @@
import json
import os
from unittest.mock import AsyncMock, patch
import httpx
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.payment.cost_calculation import CostData # noqa: E402
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
def _make_provider() -> BaseUpstreamProvider:
return BaseUpstreamProvider(base_url="http://test", api_key="test-key")
def _make_httpx_response(status_code: int = 200) -> httpx.Response:
return httpx.Response(status_code, headers={})
def _make_cost_data(total_msats: int = 5000) -> CostData:
return CostData(
base_msats=0,
input_msats=3000,
output_msats=2000,
total_msats=total_msats,
total_usd=0.00025,
input_tokens=100,
output_tokens=50,
)
# ---------------------------------------------------------------------------
# Non-streaming (chat completions)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_non_streaming_includes_cost_sats() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=5000)
response_body = {
"model": "gpt-4o",
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"total_tokens": 150,
"cost": 0.00025,
},
}
content_str = json.dumps(response_body)
httpx_response = _make_httpx_response()
with (
patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)),
patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")),
):
response = await provider.handle_x_cashu_non_streaming_response(
content_str=content_str,
response=httpx_response,
amount=10000,
unit="msat",
max_cost_for_model=10000,
mint=None,
payment_token_hash=None,
)
body = json.loads(response.body)
assert "cost_sats" in body["usage"]
assert body["usage"]["cost_sats"] == 5 # 5000 msats // 1000
@pytest.mark.asyncio
async def test_non_streaming_cost_sats_value_rounds_down() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=1999)
response_body = {"model": "gpt-4o", "usage": {"prompt_tokens": 10}}
content_str = json.dumps(response_body)
with (
patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)),
patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")),
):
response = await provider.handle_x_cashu_non_streaming_response(
content_str=content_str,
response=_make_httpx_response(),
amount=10000,
unit="msat",
max_cost_for_model=10000,
)
body = json.loads(response.body)
assert body["usage"]["cost_sats"] == 1 # 1999 // 1000
@pytest.mark.asyncio
async def test_non_streaming_preserves_existing_usage_fields() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=3000)
response_body = {
"model": "gpt-4o",
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"total_tokens": 150,
"cost": 0.00015,
},
}
with (
patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)),
patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")),
):
response = await provider.handle_x_cashu_non_streaming_response(
content_str=json.dumps(response_body),
response=_make_httpx_response(),
amount=10000,
unit="msat",
max_cost_for_model=10000,
)
body = json.loads(response.body)
usage = body["usage"]
assert usage["prompt_tokens"] == 100
assert usage["completion_tokens"] == 50
assert usage["total_tokens"] == 150
assert usage["cost"] == 0.00015
assert usage["cost_sats"] == 3
# ---------------------------------------------------------------------------
# Streaming (chat completions)
# ---------------------------------------------------------------------------
async def _collect_streaming(response: object) -> list[str]:
chunks: list[str] = []
async for chunk in response.body_iterator: # type: ignore[attr-defined]
if isinstance(chunk, bytes):
chunks.append(chunk.decode("utf-8"))
else:
chunks.append(str(chunk))
return chunks
@pytest.mark.asyncio
async def test_streaming_includes_cost_sats_in_usage_chunk() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=7000)
usage_chunk = {
"id": "chatcmpl-123",
"model": "gpt-4o",
"usage": {"prompt_tokens": 100, "completion_tokens": 50, "total_tokens": 150},
}
content_str = "\n".join([
'data: {"id":"chatcmpl-123","model":"gpt-4o","choices":[]}',
f"data: {json.dumps(usage_chunk)}",
"data: [DONE]",
])
with patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)):
response = await provider.handle_x_cashu_streaming_response(
content_str=content_str,
response=_make_httpx_response(),
amount=10000,
unit="msat",
max_cost_for_model=10000,
mint=None,
payment_token_hash=None,
)
chunks = await _collect_streaming(response)
full_output = "".join(chunks)
usage_line = next(
line for line in full_output.split("\n") if '"usage"' in line and "cost_sats" in line
)
data_json = json.loads(usage_line.lstrip("data: ").strip())
assert data_json["usage"]["cost_sats"] == 7 # 7000 // 1000
@pytest.mark.asyncio
async def test_streaming_non_usage_chunks_unmodified() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=2000)
regular_chunk = {"id": "chatcmpl-123", "model": "gpt-4o", "choices": [{"delta": {"content": "hi"}}]}
usage_chunk = {"id": "chatcmpl-123", "model": "gpt-4o", "usage": {"prompt_tokens": 10}}
content_str = "\n".join([
f"data: {json.dumps(regular_chunk)}",
f"data: {json.dumps(usage_chunk)}",
"data: [DONE]",
])
with patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)):
response = await provider.handle_x_cashu_streaming_response(
content_str=content_str,
response=_make_httpx_response(),
amount=10000,
unit="msat",
max_cost_for_model=10000,
)
chunks = await _collect_streaming(response)
lines = [
line for line in "".join(chunks).split("\n")
if line.startswith("data: ") and line != "data: [DONE]"
]
regular_line_data = json.loads(lines[0][6:])
# regular chunk should not have cost_sats injected
assert "cost_sats" not in regular_line_data.get("usage", {})

View File

@@ -14,10 +14,11 @@ import {
} from '@/lib/api/services/admin';
import { AddProviderModelDialog } from '@/components/add-provider-model-dialog';
import { BatchOverrideDialog } from '@/components/batch-override-dialog';
import { ProviderFeeScheduleModal } from '@/components/provider-fee-schedule-modal';
import { ProviderCard } from '@/components/provider-card';
import { ProviderFormDialogContent } from '@/components/provider-form-dialog-content';
import { Skeleton } from '@/components/ui/skeleton';
import { AlertCircle, Plus, Server } from 'lucide-react';
import { AlertCircle, Clock, Plus, Server } from 'lucide-react';
import { Alert, AlertDescription } from '@/components/ui/alert';
import { Dialog, DialogTrigger } from '@/components/ui/dialog';
import {
@@ -70,6 +71,10 @@ export default function ProvidersPage() {
const [batchOverrideProviderId, setBatchOverrideProviderId] = useState<
number | null
>(null);
const [feeScheduleState, setFeeScheduleState] = useState<{
open: boolean;
initialIds: number[];
}>({ open: false, initialIds: [] });
const [providerDeleteTarget, setProviderDeleteTarget] =
useState<UpstreamProvider | null>(null);
const [modelDeleteTarget, setModelDeleteTarget] = useState<{
@@ -242,6 +247,7 @@ export default function ProvidersPage() {
api_version: provider.api_version || null,
enabled: provider.enabled,
provider_fee: provider.provider_fee,
provider_fee_default: provider.provider_fee_default,
provider_settings: provider.provider_settings || {},
});
setIsEditDialogOpen(true);
@@ -254,7 +260,7 @@ export default function ProvidersPage() {
base_url: formData.base_url,
api_version: formData.api_version,
enabled: formData.enabled,
provider_fee: formData.provider_fee,
provider_fee_default: formData.provider_fee_default,
provider_settings: formData.provider_settings,
};
if (formData.api_key) {
@@ -342,6 +348,13 @@ export default function ProvidersPage() {
setBatchOverrideProviderId(providerId);
};
const handleManageFeeSchedules = (providerId?: number) => {
setFeeScheduleState({
open: true,
initialIds: providerId !== undefined ? [providerId] : [],
});
};
const availableMints = (globalSettings?.cashu_mints as string[]) || [];
return (
@@ -352,12 +365,22 @@ export default function ProvidersPage() {
title='Upstream Providers'
description='Manage your AI provider connections and credentials.'
actions={
<DialogTrigger asChild>
<Button>
<Plus className='h-4 w-4' />
Add Provider
<div className='flex gap-2'>
<Button
variant='outline'
onClick={() => handleManageFeeSchedules()}
disabled={providers.length === 0}
>
<Clock className='h-4 w-4' />
Fee Schedules
</Button>
</DialogTrigger>
<DialogTrigger asChild>
<Button>
<Plus className='h-4 w-4' />
Add Provider
</Button>
</DialogTrigger>
</div>
}
/>
<ProviderFormDialogContent
@@ -437,6 +460,9 @@ export default function ProvidersPage() {
onEditProvider={() => handleEdit(provider)}
onDeleteProvider={() => setProviderDeleteTarget(provider)}
onBatchOverride={() => handleBatchOverride(provider.id)}
onManageFeeSchedules={() =>
handleManageFeeSchedules(provider.id)
}
onAddModel={() => handleAddModel(provider.id)}
onEditModel={(model) => handleEditModel(provider.id, model)}
onDeleteModel={(modelId) =>
@@ -562,6 +588,16 @@ export default function ProvidersPage() {
}}
/>
)}
<ProviderFeeScheduleModal
providers={providers}
initialSelectedIds={feeScheduleState.initialIds}
isOpen={feeScheduleState.open}
onClose={() => setFeeScheduleState({ open: false, initialIds: [] })}
onSuccess={() => {
queryClient.invalidateQueries({ queryKey: ['upstream-providers'] });
}}
/>
</AppPageShell>
);
}

View File

@@ -4,6 +4,7 @@ import * as React from 'react';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
import { ServerConfigSettings } from '@/components/settings/server-config-settings';
import { AdminSettings } from '@/components/settings/admin-settings';
import { CliTokensSettings } from '@/components/settings/cli-tokens-settings';
import { AppPageShell } from '@/components/app-page-shell';
import { PageHeader } from '@/components/page-header';
@@ -19,6 +20,7 @@ export default function SettingsPage() {
<TabsList variant='line' className='mb-4 w-full'>
<TabsTrigger value='admin'>Admin Settings</TabsTrigger>
<TabsTrigger value='server'>Server Config</TabsTrigger>
<TabsTrigger value='cli-tokens'>CLI Tokens</TabsTrigger>
</TabsList>
<TabsContent value='server'>
<ServerConfigSettings />
@@ -26,6 +28,9 @@ export default function SettingsPage() {
<TabsContent value='admin'>
<AdminSettings />
</TabsContent>
<TabsContent value='cli-tokens'>
<CliTokensSettings />
</TabsContent>
</Tabs>
</div>
</AppPageShell>

View File

@@ -1,7 +1,7 @@
'use client';
import { useState, useEffect } from 'react';
import { useQuery } from '@tanstack/react-query';
import { useQuery, keepPreviousData } from '@tanstack/react-query';
import { AppPageShell } from '@/components/app-page-shell';
import { PageHeader } from '@/components/page-header';
import {
@@ -22,6 +22,7 @@ import {
SelectValue,
} from '@/components/ui/select';
import { Badge } from '@/components/ui/badge';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
import {
Table,
TableBody,
@@ -30,7 +31,7 @@ import {
TableHeader,
TableRow,
} from '@/components/ui/table';
import { ScrollArea } from '@/components/ui/scroll-area';
import { ScrollArea, ScrollBar } from '@/components/ui/scroll-area';
import { Skeleton } from '@/components/ui/skeleton';
import {
Empty,
@@ -47,6 +48,10 @@ import {
Copy,
Check,
Receipt,
Key,
Zap,
ChevronLeft,
ChevronRight,
} from 'lucide-react';
import { AdminService, type Transaction } from '@/lib/api/services/admin';
import { format } from 'date-fns';
@@ -54,6 +59,147 @@ import { toast } from 'sonner';
const STORAGE_KEY = 'routstr-transaction-filters';
function TransactionTable({
transactions,
copiedId,
onCopy,
getStatusBadge,
}: {
transactions: Transaction[];
copiedId: string | null;
onCopy: (text: string, id: string) => void;
getStatusBadge: (tx: Transaction) => React.ReactNode;
}) {
if (transactions.length === 0) {
return (
<Empty className='py-8'>
<EmptyHeader>
<EmptyMedia variant='icon'>
<Receipt className='h-4 w-4' />
</EmptyMedia>
<EmptyTitle>No transactions found</EmptyTitle>
<EmptyDescription>
Try adjusting your filters or check back later.
</EmptyDescription>
</EmptyHeader>
</Empty>
);
}
return (
<ScrollArea className='h-[55svh] min-h-[420px] w-full sm:h-[600px]'>
<div className='min-w-[800px]'>
<Table>
<TableHeader>
<TableRow>
<TableHead>Type</TableHead>
<TableHead>Amount</TableHead>
<TableHead>Status</TableHead>
<TableHead>API Key</TableHead>
<TableHead>Request ID</TableHead>
<TableHead>Mint</TableHead>
<TableHead>Date</TableHead>
<TableHead className='text-right'>Actions</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{transactions.map((tx) => (
<TableRow key={tx.id}>
<TableCell>
<div className='flex items-center gap-2'>
{tx.type === 'in' ? (
<ArrowDownLeft className='h-4 w-4 text-green-500' />
) : (
<ArrowUpRight className='h-4 w-4 text-blue-500' />
)}
<span className='capitalize'>{tx.type}</span>
</div>
</TableCell>
<TableCell className='font-mono'>
{tx.amount} {tx.unit}
</TableCell>
<TableCell>{getStatusBadge(tx)}</TableCell>
<TableCell>
{tx.api_key_hashed_key ? (
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[120px] truncate font-mono'>
{tx.api_key_hashed_key.slice(0, 12)}...
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() =>
onCopy(tx.api_key_hashed_key!, tx.id + '-apikey')
}
>
{copiedId === tx.id + '-apikey' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
) : (
<span className='text-muted-foreground text-xs'></span>
)}
</TableCell>
<TableCell>
{tx.request_id ? (
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[150px] truncate font-mono'>
{tx.request_id}
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() => onCopy(tx.request_id!, tx.id + '-req')}
>
{copiedId === tx.id + '-req' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
) : (
<span className='text-muted-foreground text-xs'></span>
)}
</TableCell>
<TableCell>
<div className='flex max-w-[150px] items-center gap-1 truncate text-xs'>
<span className='truncate'>{tx.mint_url}</span>
</div>
</TableCell>
<TableCell className='text-xs whitespace-nowrap'>
{format(tx.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')}
</TableCell>
<TableCell className='text-right'>
<Button
variant='ghost'
size='icon'
className='h-8 w-8'
onClick={() => onCopy(tx.token, tx.id + '-token')}
title='Copy Token'
>
{copiedId === tx.id + '-token' ? (
<Check className='h-4 w-4' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
<ScrollBar orientation='horizontal' />
</ScrollArea>
);
}
export default function TransactionsPage() {
const [search, setSearch] = useState('');
const [type, setType] = useState<string>('all');
@@ -81,21 +227,63 @@ export default function TransactionsPage() {
localStorage.setItem(STORAGE_KEY, JSON.stringify(filters));
}, [search, type, status]);
const { data, isLoading, refetch, isRefetching } = useQuery({
queryKey: ['transactions', type, status, search],
const PAGE_SIZE = 50;
const [activeTab, setActiveTab] = useState<string>('x-cashu');
const [xcashuPage, setXcashuPage] = useState(0);
const [apikeyPage, setApikeyPage] = useState(0);
const typeParam = type === 'all' ? undefined : type;
const statusParam = status === 'all' ? undefined : status;
const searchParam = search || undefined;
const xcashuQuery = useQuery({
queryKey: [
'transactions',
'x-cashu',
typeParam,
statusParam,
searchParam,
xcashuPage,
],
queryFn: () =>
AdminService.getTransactions(
type === 'all' ? undefined : type,
status === 'all' ? undefined : status,
search || undefined,
100
typeParam,
statusParam,
searchParam,
'x-cashu',
PAGE_SIZE,
xcashuPage * PAGE_SIZE
),
placeholderData: keepPreviousData,
});
const apikeyQuery = useQuery({
queryKey: [
'transactions',
'apikey',
typeParam,
statusParam,
searchParam,
apikeyPage,
],
queryFn: () =>
AdminService.getTransactions(
typeParam,
statusParam,
searchParam,
'apikey',
PAGE_SIZE,
apikeyPage * PAGE_SIZE
),
placeholderData: keepPreviousData,
});
const handleClearFilters = () => {
setSearch('');
setType('all');
setStatus('all');
setXcashuPage(0);
setApikeyPage(0);
};
const copyToClipboard = (text: string, id: string) => {
@@ -145,15 +333,91 @@ export default function TransactionsPage() {
.filter(Boolean)
.join(' • ');
// Reset pages when filters change
useEffect(() => {
setXcashuPage(0);
setApikeyPage(0);
}, [type, status, search]);
const isRefetching = xcashuQuery.isRefetching || apikeyQuery.isRefetching;
const renderCardContent = (
query: typeof xcashuQuery,
page: number,
setPage: (p: number) => void
) => {
if (query.isLoading) {
return (
<div className='space-y-2'>
{Array.from({ length: 8 }).map((_, index) => (
<Skeleton
key={`tx-loading-${index}`}
className='h-16 w-full rounded-lg'
/>
))}
</div>
);
}
const transactions = query.data?.transactions ?? [];
const total = query.data?.total ?? 0;
const totalPages = Math.ceil(total / PAGE_SIZE);
return (
<>
{totalPages > 1 && (
<div className='flex flex-col gap-2 border-b pb-3 sm:flex-row sm:items-center sm:justify-between'>
<span className='text-muted-foreground text-xs sm:text-sm'>
{page * PAGE_SIZE + 1}{Math.min((page + 1) * PAGE_SIZE, total)}{' '}
of {total}
</span>
<div className='flex items-center gap-2'>
<Button
variant='outline'
size='sm'
disabled={page === 0}
onClick={() => setPage(page - 1)}
>
<ChevronLeft className='h-4 w-4' />
<span className='hidden sm:inline'>Previous</span>
</Button>
<span className='text-xs sm:text-sm'>
{page + 1} / {totalPages}
</span>
<Button
variant='outline'
size='sm'
disabled={page >= totalPages - 1}
onClick={() => setPage(page + 1)}
>
<span className='hidden sm:inline'>Next</span>
<ChevronRight className='h-4 w-4' />
</Button>
</div>
</div>
)}
<TransactionTable
transactions={transactions}
copiedId={copiedId}
onCopy={copyToClipboard}
getStatusBadge={getStatusBadge}
/>
</>
);
};
return (
<AppPageShell contentClassName='mx-auto w-full max-w-5xl overflow-x-hidden'>
<div className='space-y-6'>
<PageHeader
title='X-Cashu Transactions'
description='View all incoming and outgoing X-Cashu token transactions.'
title='Cashu Transactions'
description='View all incoming and outgoing Cashu token transactions.'
actions={
<Button
onClick={() => refetch()}
onClick={() => {
xcashuQuery.refetch();
apikeyQuery.refetch();
}}
variant='outline'
size='sm'
disabled={isRefetching}
@@ -181,7 +445,7 @@ export default function TransactionsPage() {
<Search className='text-muted-foreground absolute top-2.5 left-2.5 h-4 w-4' />
<Input
id='search'
placeholder='Search by ID, token or request ID...'
placeholder='Search by ID, token, request ID or key hash...'
className='pl-8'
value={search}
onChange={(e) => setSearch(e.target.value)}
@@ -228,138 +492,68 @@ export default function TransactionsPage() {
</CardContent>
</Card>
<Card>
<CardHeader>
<div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'>
<CardTitle>Transaction History</CardTitle>
{data && (
<Badge variant='secondary'>
{data.transactions.length} entries
<Tabs
defaultValue='x-cashu'
value={activeTab}
onValueChange={setActiveTab}
>
<TabsList className='mb-4'>
<TabsTrigger value='x-cashu' className='flex items-center gap-2'>
<Zap className='h-4 w-4' />
X-Cashu
{xcashuQuery.data && (
<Badge variant='secondary' className='ml-1'>
{xcashuQuery.data.total}
</Badge>
)}
</div>
{hasActiveFilters && (
<CardDescription>
Showing transactions filtered by {activeFilterDescription}
</CardDescription>
)}
</CardHeader>
<CardContent className='overflow-hidden'>
{isLoading ? (
<div className='space-y-2'>
{Array.from({ length: 8 }).map((_, index) => (
<Skeleton
key={`tx-loading-${index}`}
className='h-16 w-full rounded-lg'
/>
))}
</div>
) : data?.transactions && data.transactions.length > 0 ? (
<ScrollArea className='h-[55svh] min-h-[420px] w-full sm:h-[600px]'>
<Table>
<TableHeader>
<TableRow>
<TableHead>Type</TableHead>
<TableHead>Amount</TableHead>
<TableHead>Status</TableHead>
<TableHead>Request ID</TableHead>
<TableHead>Mint</TableHead>
<TableHead>Date</TableHead>
<TableHead className='text-right'>Actions</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{data.transactions.map((tx) => (
<TableRow key={tx.id}>
<TableCell>
<div className='flex items-center gap-2'>
{tx.type === 'in' ? (
<ArrowDownLeft className='h-4 w-4 text-green-500' />
) : (
<ArrowUpRight className='h-4 w-4 text-blue-500' />
)}
<span className='capitalize'>{tx.type}</span>
</div>
</TableCell>
<TableCell className='font-mono'>
{tx.amount} {tx.unit}
</TableCell>
<TableCell>{getStatusBadge(tx)}</TableCell>
<TableCell>
{tx.request_id ? (
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[150px] truncate font-mono'>
{tx.request_id}
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() =>
copyToClipboard(
tx.request_id!,
tx.id + '-req'
)
}
>
{copiedId === tx.id + '-req' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
) : (
<span className='text-muted-foreground text-xs'>
</span>
)}
</TableCell>
<TableCell>
<div className='flex max-w-[150px] items-center gap-1 truncate text-xs'>
<span className='truncate'>{tx.mint_url}</span>
</div>
</TableCell>
<TableCell className='text-xs whitespace-nowrap'>
{format(tx.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')}
</TableCell>
<TableCell className='text-right'>
<Button
variant='ghost'
size='icon'
className='h-8 w-8'
onClick={() =>
copyToClipboard(tx.token, tx.id + '-token')
}
title='Copy Token'
>
{copiedId === tx.id + '-token' ? (
<Check className='h-4 w-4' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</ScrollArea>
) : (
<Empty className='py-8'>
<EmptyHeader>
<EmptyMedia variant='icon'>
<Receipt className='h-4 w-4' />
</EmptyMedia>
<EmptyTitle>No transactions found</EmptyTitle>
<EmptyDescription>
Try adjusting your filters or check back later.
</EmptyDescription>
</EmptyHeader>
</Empty>
)}
</CardContent>
</Card>
</TabsTrigger>
<TabsTrigger value='apikey' className='flex items-center gap-2'>
<Key className='h-4 w-4' />
API Key
{apikeyQuery.data && (
<Badge variant='secondary' className='ml-1'>
{apikeyQuery.data.total}
</Badge>
)}
</TabsTrigger>
</TabsList>
<TabsContent value='x-cashu'>
<Card>
<CardHeader>
<div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'>
<CardTitle>X-Cashu Transaction History</CardTitle>
{hasActiveFilters && (
<CardDescription>
Filtered by {activeFilterDescription}
</CardDescription>
)}
</div>
</CardHeader>
<CardContent className='overflow-hidden'>
{renderCardContent(xcashuQuery, xcashuPage, setXcashuPage)}
</CardContent>
</Card>
</TabsContent>
<TabsContent value='apikey'>
<Card>
<CardHeader>
<div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'>
<CardTitle>API Key Transaction History</CardTitle>
{hasActiveFilters && (
<CardDescription>
Filtered by {activeFilterDescription}
</CardDescription>
)}
</div>
</CardHeader>
<CardContent className='overflow-hidden'>
{renderCardContent(apikeyQuery, apikeyPage, setApikeyPage)}
</CardContent>
</Card>
</TabsContent>
</Tabs>
</div>
</AppPageShell>
);

View File

@@ -1,6 +1,6 @@
'use client';
import React, { useEffect, useMemo, useState } from 'react';
import React, { useCallback, useEffect, useMemo, useState } from 'react';
import { useForm } from 'react-hook-form';
import { z } from 'zod';
import { zodResolver } from '@hookform/resolvers/zod';
@@ -39,7 +39,7 @@ import {
FormMessage,
} from '@/components/ui/form';
import { Switch } from '@/components/ui/switch';
import { Loader2, Plus } from 'lucide-react';
import { Check, Copy, Loader2, Plus } from 'lucide-react';
import { toast } from 'sonner';
import { AdminService, type AdminModel } from '@/lib/api/services/admin';
@@ -64,6 +64,7 @@ const FormSchema = z.object({
instruct_type: z.string().default(''),
canonical_slug: z.string().default(''),
alias_ids_raw: z.string().default(''),
forwarded_model_id: z.string().default(''),
upstream_provider_id: z.string().default(''),
input_cost: z.coerce.number().min(0).default(0),
output_cost: z.coerce.number().min(0).default(0),
@@ -104,6 +105,7 @@ export function AddProviderModelDialog({
const [isPresetOpen, setIsPresetOpen] = useState(false);
const [selectedPresetLabel, setSelectedPresetLabel] =
useState('Select a preset');
const [forwardedModelIdCopied, setForwardedModelIdCopied] = useState(false);
const form = useForm<FormData>({
resolver: zodResolver(FormSchema) as never,
@@ -119,6 +121,7 @@ export function AddProviderModelDialog({
instruct_type: '',
canonical_slug: '',
alias_ids_raw: '',
forwarded_model_id: '',
upstream_provider_id: '',
input_cost: 0,
output_cost: 0,
@@ -180,6 +183,7 @@ export function AddProviderModelDialog({
: '',
canonical_slug: initialData.canonical_slug || '',
alias_ids_raw: listToString(initialData.alias_ids),
forwarded_model_id: initialData.forwarded_model_id || initialData.id,
upstream_provider_id:
typeof initialData.upstream_provider_id === 'string'
? initialData.upstream_provider_id
@@ -223,6 +227,7 @@ export function AddProviderModelDialog({
instruct_type: '',
canonical_slug: '',
alias_ids_raw: '',
forwarded_model_id: '',
upstream_provider_id: '',
input_cost: 0,
output_cost: 0,
@@ -280,6 +285,7 @@ export function AddProviderModelDialog({
);
form.setValue('canonical_slug', model.canonical_slug || '');
form.setValue('alias_ids_raw', listToString(model.alias_ids));
form.setValue('forwarded_model_id', model.forwarded_model_id || model.id);
form.setValue(
'upstream_provider_id',
typeof model.upstream_provider_id === 'string'
@@ -385,6 +391,7 @@ export function AddProviderModelDialog({
canonical_slug: data.canonical_slug?.trim() || null,
alias_ids: listFromString(data.alias_ids_raw || ''),
enabled: data.enabled,
forwarded_model_id: data.forwarded_model_id?.trim() || data.id,
};
if (isEdit) {
@@ -520,6 +527,53 @@ export function AddProviderModelDialog({
)}
/>
<FormField
control={form.control}
name='forwarded_model_id'
render={({ field }) => {
const handleCopy = () => {
const value = field.value || form.getValues('id');
if (!value) return;
navigator.clipboard.writeText(value);
setForwardedModelIdCopied(true);
setTimeout(() => setForwardedModelIdCopied(false), 1500);
};
return (
<FormItem>
<FormLabel>Client Alias ID</FormLabel>
<FormControl>
<div className='flex gap-2'>
<Input
placeholder={
form.watch('id') || 'e.g., openai/gpt-4o'
}
{...field}
/>
<Button
type='button'
variant='outline'
size='icon'
onClick={handleCopy}
title='Copy model ID'
>
{forwardedModelIdCopied ? (
<Check className='h-4 w-4 text-green-500' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</div>
</FormControl>
<FormDescription>
Alternate ID that clients can use to reference this
model. Defaults to the model&apos;s own ID.
</FormDescription>
<FormMessage />
</FormItem>
);
}}
/>
<FormField
control={form.control}
name='name'

View File

@@ -471,7 +471,11 @@ export function ApiEndpointTester({ models }: ApiEndpointTesterProps) {
testEndpointMutation.mutate(requestData);
};
const enabledModels = models.filter((model) => model.isEnabled);
const enabledModels = Array.from(
new Map(
models.filter((model) => model.isEnabled).map((m) => [m.id, m])
).values()
);
const credentials = selectedModel ? getModelCredentials(selectedModel) : null;
const endpointUrl = credentials
? buildEndpointUrl(

View File

@@ -1,577 +0,0 @@
'use client';
import React, { useState, useEffect, useCallback } from 'react';
import { useForm } from 'react-hook-form';
import { zodResolver } from '@hookform/resolvers/zod';
import { z } from 'zod';
import { type Model } from '@/lib/api/schemas/models';
import { AdminService, type AdminModel } from '@/lib/api/services/admin';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Textarea } from '@/components/ui/textarea';
import {
Dialog,
DialogContent,
DialogDescription,
DialogHeader,
DialogTitle,
} from '@/components/ui/dialog';
import {
Form,
FormControl,
FormDescription,
FormField,
FormItem,
FormLabel,
FormMessage,
} from '@/components/ui/form';
import { Edit3, Loader2 } from 'lucide-react';
import { toast } from 'sonner';
import { Switch } from '@/components/ui/switch';
const EditModelFormSchema = z.object({
name: z.string().min(1, 'Name is required'),
description: z.string().optional(),
context_length: z.number().min(0),
prompt: z.number().min(0),
completion: z.number().min(0),
enabled: z.boolean(),
});
type EditModelFormData = z.infer<typeof EditModelFormSchema>;
const roundToFiveDecimals = (value: number | undefined | null): number => {
if (value === undefined || value === null || isNaN(value)) {
return 0;
}
return Math.round(value * 100000) / 100000;
};
const toNumber = (value: unknown, fallback = 0): number => {
if (typeof value === 'number' && Number.isFinite(value)) {
return value;
}
if (typeof value === 'string') {
const parsed = Number(value);
if (Number.isFinite(parsed)) {
return parsed;
}
}
return fallback;
};
const toStringArray = (value: unknown, fallback: string[]): string[] => {
if (!Array.isArray(value)) {
return fallback;
}
const filtered = value.filter(
(item): item is string => typeof item === 'string'
);
return filtered.length > 0 ? filtered : fallback;
};
interface EditModelFormProps {
model: Model;
providerId?: number;
onModelUpdate?: () => void;
onCancel?: () => void;
isOpen: boolean;
}
interface AdminModelData {
id: string;
name: string;
description?: string;
created: number;
context_length: number;
architecture: {
modality: string;
input_modalities: string[];
output_modalities: string[];
tokenizer: string;
instruct_type: string | null;
};
pricing: {
prompt: number;
completion: number;
request: number;
image: number;
web_search: number;
internal_reasoning: number;
};
per_request_limits: null | undefined;
top_provider: null | undefined;
upstream_provider_id: number;
enabled: boolean;
}
const normalizeAdminModelData = (
adminModel: AdminModel,
fallbackModel: Model,
providerId: number
): AdminModelData => {
const pricingRecord =
adminModel.pricing && typeof adminModel.pricing === 'object'
? (adminModel.pricing as Record<string, unknown>)
: {};
const architectureRecord =
adminModel.architecture && typeof adminModel.architecture === 'object'
? (adminModel.architecture as Record<string, unknown>)
: {};
return {
id: adminModel.id,
name: adminModel.name,
description: adminModel.description || '',
created: toNumber(adminModel.created, Math.floor(Date.now() / 1000)),
context_length: Math.max(
0,
Math.trunc(
toNumber(adminModel.context_length, fallbackModel.contextLength || 4096)
)
),
architecture: {
modality:
typeof architectureRecord.modality === 'string'
? architectureRecord.modality
: fallbackModel.modelType || 'text',
input_modalities: toStringArray(architectureRecord.input_modalities, [
fallbackModel.modelType || 'text',
]),
output_modalities: toStringArray(architectureRecord.output_modalities, [
fallbackModel.modelType || 'text',
]),
tokenizer:
typeof architectureRecord.tokenizer === 'string'
? architectureRecord.tokenizer
: '',
instruct_type:
typeof architectureRecord.instruct_type === 'string'
? architectureRecord.instruct_type
: null,
},
pricing: {
prompt: roundToFiveDecimals(
toNumber(pricingRecord.prompt, fallbackModel.input_cost)
),
completion: roundToFiveDecimals(
toNumber(pricingRecord.completion, fallbackModel.output_cost)
),
request: toNumber(pricingRecord.request, 0),
image: toNumber(pricingRecord.image, 0),
web_search: toNumber(pricingRecord.web_search, 0),
internal_reasoning: toNumber(pricingRecord.internal_reasoning, 0),
},
per_request_limits:
adminModel.per_request_limits === null ||
adminModel.per_request_limits === undefined
? adminModel.per_request_limits
: null,
top_provider:
adminModel.top_provider === null || adminModel.top_provider === undefined
? adminModel.top_provider
: null,
upstream_provider_id:
typeof adminModel.upstream_provider_id === 'number'
? adminModel.upstream_provider_id
: providerId,
enabled: adminModel.enabled !== false,
};
};
export function EditModelForm({
model,
providerId,
onModelUpdate,
onCancel,
isOpen,
}: EditModelFormProps) {
const [isSubmitting, setIsSubmitting] = useState(false);
const [adminModelData, setAdminModelData] = useState<AdminModelData | null>(
null
);
const [isNewOverride, setIsNewOverride] = useState(false);
const form = useForm<EditModelFormData>({
resolver: zodResolver(EditModelFormSchema),
defaultValues: {
name: model.name,
description: model.description || '',
context_length: model.contextLength || 4096,
prompt: roundToFiveDecimals(model.input_cost),
completion: roundToFiveDecimals(model.output_cost),
enabled: model.isEnabled !== false,
},
});
const loadAdminModel = useCallback(async () => {
if (!providerId) {
console.error('loadAdminModel called without providerId');
return;
}
try {
const adminModel = await AdminService.getProviderModel(
providerId,
model.id
);
const normalizedAdminModel = normalizeAdminModelData(
adminModel,
model,
providerId
);
setAdminModelData(normalizedAdminModel);
setIsNewOverride(false);
form.reset({
name: normalizedAdminModel.name,
description: normalizedAdminModel.description || '',
context_length: normalizedAdminModel.context_length,
prompt: normalizedAdminModel.pricing.prompt,
completion: normalizedAdminModel.pricing.completion,
enabled: normalizedAdminModel.enabled !== false,
});
} catch {
setIsNewOverride(true);
setAdminModelData({
id: model.full_name,
name: model.name,
description: model.description || '',
created: Math.floor(Date.now() / 1000),
context_length: model.contextLength || 4096,
architecture: {
modality: model.modelType || 'text',
input_modalities: [model.modelType || 'text'],
output_modalities: [model.modelType || 'text'],
tokenizer: '',
instruct_type: null,
},
pricing: {
prompt: roundToFiveDecimals(model.input_cost),
completion: roundToFiveDecimals(model.output_cost),
request: 0,
image: 0,
web_search: 0,
internal_reasoning: 0,
},
per_request_limits: null,
top_provider: null,
upstream_provider_id: providerId,
enabled: model.isEnabled !== false,
});
form.reset({
name: model.name,
description: model.description || '',
context_length: model.contextLength || 4096,
prompt: roundToFiveDecimals(model.input_cost),
completion: roundToFiveDecimals(model.output_cost),
enabled: model.isEnabled !== false,
});
}
}, [providerId, model, form]);
useEffect(() => {
if (isOpen && providerId) {
loadAdminModel();
} else if (isOpen && !providerId) {
console.error('EditModelForm opened without providerId', {
model,
providerId,
});
toast.error('Missing provider information for this model');
}
}, [isOpen, providerId, model, loadAdminModel]);
const onSubmit = async (data: EditModelFormData) => {
if (!providerId) {
console.error('onSubmit called without providerId', {
model,
providerId,
});
toast.error('Missing provider ID - cannot update model');
return;
}
if (!adminModelData) {
console.error('onSubmit called without adminModelData', {
model,
providerId,
adminModelData,
});
toast.error('Model data not loaded - please try reopening the form');
return;
}
setIsSubmitting(true);
try {
const payload = {
id: adminModelData.id,
name: data.name,
description: data.description || '',
created: adminModelData.created || Math.floor(Date.now() / 1000),
context_length: data.context_length,
architecture: adminModelData.architecture || {
modality: 'text',
input_modalities: ['text'],
output_modalities: ['text'],
tokenizer: '',
instruct_type: null,
},
pricing: {
prompt: roundToFiveDecimals(data.prompt),
completion: roundToFiveDecimals(data.completion),
request: 0,
image: 0,
web_search: 0,
internal_reasoning: 0,
},
per_request_limits: adminModelData.per_request_limits,
top_provider: adminModelData.top_provider,
upstream_provider_id: providerId,
enabled: data.enabled,
};
if (isNewOverride) {
await AdminService.createProviderModel(providerId, payload);
toast.success('Model override created successfully!');
} else {
await AdminService.updateProviderModel(
providerId,
adminModelData.id,
payload
);
toast.success('Model updated successfully!');
}
onModelUpdate?.();
onCancel?.();
} catch (error) {
const action = isNewOverride ? 'create' : 'update';
toast.error(`Failed to ${action} model. Please try again.`);
console.error(`Error ${action}ing model:`, error);
} finally {
setIsSubmitting(false);
}
};
const handleClose = () => {
if (!isSubmitting) {
onCancel?.();
}
};
return (
<Dialog open={isOpen} onOpenChange={handleClose}>
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[600px]'>
<DialogHeader>
<DialogTitle className='flex items-center gap-2'>
<Edit3 className='h-5 w-5' />
{isNewOverride ? 'Create Model Override' : 'Edit Model Override'}
</DialogTitle>
<DialogDescription>
{isNewOverride
? `Create an override for &quot;${model.name}&quot;`
: `Update the model override for &quot;${model.name}&quot;`}
</DialogDescription>
</DialogHeader>
<Form {...form}>
<form onSubmit={form.handleSubmit(onSubmit)} className='space-y-4'>
<div className='grid grid-cols-1 gap-4 sm:grid-cols-2'>
<FormField
control={form.control}
name='name'
render={({ field }) => (
<FormItem>
<FormLabel>Display Name *</FormLabel>
<FormControl>
<Input
placeholder='e.g., GPT-4'
{...field}
className='w-full'
/>
</FormControl>
<FormDescription>
Custom display name for the model
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
<FormField
control={form.control}
name='context_length'
render={({ field }) => (
<FormItem>
<FormLabel>Context Length *</FormLabel>
<FormControl>
<Input
type='number'
min='0'
placeholder='4096'
value={field.value ?? ''}
onChange={(e) => {
const value = e.target.value;
field.onChange(
value === '' ? 0 : parseInt(value, 10) || 0
);
}}
onBlur={(e) => {
const value = parseInt(e.target.value, 10);
field.onChange(
Number.isNaN(value) ? 0 : Math.max(0, value)
);
}}
className='w-full'
/>
</FormControl>
<FormDescription>
Maximum context window size
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
</div>
<FormField
control={form.control}
name='description'
render={({ field }) => (
<FormItem>
<FormLabel>Description</FormLabel>
<FormControl>
<Textarea
placeholder='Brief description of the model...'
{...field}
rows={3}
className='w-full'
/>
</FormControl>
<FormDescription>
Optional description or notes about the model
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
<div className='grid grid-cols-1 gap-4 sm:grid-cols-2'>
<FormField
control={form.control}
name='prompt'
render={({ field }) => (
<FormItem>
<FormLabel>Input Cost (per 1M tokens) *</FormLabel>
<FormControl>
<Input
type='number'
step='0.00001'
min='0'
placeholder='5.00000'
value={field.value ?? ''}
onChange={(e) => {
const value = e.target.value;
field.onChange(value === '' ? 0 : parseFloat(value));
}}
onBlur={(e) => {
const value = parseFloat(e.target.value);
field.onChange(roundToFiveDecimals(value));
}}
className='w-full'
/>
</FormControl>
<FormDescription>
Cost in USD per 1,000,000 input tokens (max 5 decimals)
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
<FormField
control={form.control}
name='completion'
render={({ field }) => (
<FormItem>
<FormLabel>Output Cost (per 1M tokens) *</FormLabel>
<FormControl>
<Input
type='number'
step='0.00001'
min='0'
placeholder='15.00000'
value={field.value ?? ''}
onChange={(e) => {
const value = e.target.value;
field.onChange(value === '' ? 0 : parseFloat(value));
}}
onBlur={(e) => {
const value = parseFloat(e.target.value);
field.onChange(roundToFiveDecimals(value));
}}
className='w-full'
/>
</FormControl>
<FormDescription>
Cost in USD per 1,000,000 output tokens (max 5 decimals)
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
</div>
<FormField
control={form.control}
name='enabled'
render={({ field }) => (
<FormItem className='flex flex-row items-center justify-between rounded-lg border p-4'>
<div className='space-y-0.5'>
<FormLabel className='text-base'>Model Enabled</FormLabel>
<FormDescription>
Enable or disable this model override
</FormDescription>
</div>
<FormControl>
<Switch
checked={field.value}
onCheckedChange={field.onChange}
/>
</FormControl>
</FormItem>
)}
/>
<div className='flex justify-end gap-2 pt-4'>
<Button
type='button'
variant='outline'
onClick={handleClose}
disabled={isSubmitting}
>
Cancel
</Button>
<Button type='submit' disabled={isSubmitting}>
{isSubmitting ? (
<>
<Loader2 className='mr-2 h-4 w-4 animate-spin' />
{isNewOverride ? 'Creating...' : 'Updating...'}
</>
) : isNewOverride ? (
'Create Override'
) : (
'Update Model'
)}
</Button>
</div>
</form>
</Form>
</DialogContent>
</Dialog>
);
}

View File

@@ -197,7 +197,11 @@ export function ModelTester({ models }: ModelTesterProps) {
testModelMutation.mutate(request);
};
const enabledModels = models.filter((model) => model.isEnabled);
const enabledModels = Array.from(
new Map(
models.filter((model) => model.isEnabled).map((m) => [m.id, m])
).values()
);
const credentials = selectedModel ? getModelCredentials(selectedModel) : null;
return (

View File

@@ -124,12 +124,14 @@ export function ModelsPage() {
>
Basic Testing
</TabsTrigger>
{/*
<TabsTrigger
value='test-api'
className='h-9 snap-start px-2 text-[13px] sm:h-10 sm:px-2.5 sm:text-sm'
>
API Endpoints
</TabsTrigger>
*/}
</TabsList>
<TabsContent value='manage' className='mt-0'>

View File

@@ -20,6 +20,7 @@ import {
Trash2,
Key,
RotateCcw,
Clock,
} from 'lucide-react';
import { ProviderBalance } from '@/components/provider-balance';
import { ProviderModelsPanel } from '@/components/provider-models-panel';
@@ -54,6 +55,7 @@ interface ProviderCardProps {
onDeleteModel: (modelId: string) => void;
onOverrideModel: (model: AdminModel) => void;
onUpdateApiKey: (newKey: string) => void;
onManageFeeSchedules: () => void;
availableMints: string[];
}
@@ -74,6 +76,7 @@ export function ProviderCard({
onDeleteModel,
onOverrideModel,
onUpdateApiKey,
onManageFeeSchedules,
}: ProviderCardProps) {
const queryClient = useQueryClient();
const [isKeyModalOpen, setIsKeyModalOpen] = useState(false);
@@ -113,6 +116,14 @@ export function ProviderCard({
>
{provider.enabled ? 'Enabled' : 'Disabled'}
</Badge>
<Badge variant='outline' className='w-fit'>
Fee: {provider.provider_fee}x
{provider.provider_fee !== provider.provider_fee_default && (
<span className='text-muted-foreground ml-1 font-normal'>
(default: {provider.provider_fee_default}x)
</span>
)}
</Badge>
</div>
<CardDescription className='break-all'>
{provider.base_url}
@@ -189,6 +200,22 @@ export function ProviderCard({
)}
</Button>
<Button
variant='outline'
size='sm'
onClick={onManageFeeSchedules}
className='justify-center gap-1.5'
title='Manage fee schedules'
>
<Clock className='h-4 w-4' />
<span>Fees</span>
{(provider.provider_fee_schedules?.length ?? 0) > 0 && (
<Badge variant='secondary' className='ml-0.5 h-4 px-1 text-xs'>
{provider.provider_fee_schedules!.length}
</Badge>
)}
</Button>
<Button
variant='outline'
size='sm'

View File

@@ -0,0 +1,500 @@
'use client';
import { useEffect, useState } from 'react';
import { useQueryClient } from '@tanstack/react-query';
import { toast } from 'sonner';
import { Plus, Trash2 } from 'lucide-react';
import {
Dialog,
DialogContent,
DialogDescription,
DialogFooter,
DialogHeader,
DialogTitle,
} from '@/components/ui/dialog';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Label } from '@/components/ui/label';
import { Badge } from '@/components/ui/badge';
import { Checkbox } from '@/components/ui/checkbox';
import {
AdminService,
FeeTimeRange,
UpstreamProvider,
} from '@/lib/api/services/admin';
// ---------------------------------------------------------------------------
// Types
// ---------------------------------------------------------------------------
export interface ProviderFeeScheduleModalProps {
providers: UpstreamProvider[];
/** Pre-selected provider IDs (e.g. clicked from a card). Empty = all selected. */
initialSelectedIds?: number[];
isOpen: boolean;
onClose: () => void;
onSuccess: () => void;
}
interface RangeRow extends FeeTimeRange {
_id: number;
}
// ---------------------------------------------------------------------------
// Overlap helpers (mirrored from backend logic)
// ---------------------------------------------------------------------------
function _toMinutes(t: string): number {
const [h, m] = t.split(':').map(Number);
return h * 60 + m;
}
function _rangeIntervals(start: string, end: string): Array<[number, number]> {
const s = _toMinutes(start);
const e = _toMinutes(end);
if (s < e) return [[s, e]];
if (s > e)
return [
[s, 1440],
[0, e],
];
return [[0, 1440]];
}
function _intervalsOverlap(a: [number, number], b: [number, number]): boolean {
return a[0] < b[1] && b[0] < a[1];
}
function findOverlappingIds(rows: RangeRow[]): Set<number> {
const overlapping = new Set<number>();
for (let i = 0; i < rows.length; i++) {
for (let j = i + 1; j < rows.length; j++) {
const a = rows[i];
const b = rows[j];
if (!a.start_time || !a.end_time || !b.start_time || !b.end_time)
continue;
for (const ia of _rangeIntervals(a.start_time, a.end_time)) {
for (const ib of _rangeIntervals(b.start_time, b.end_time)) {
if (_intervalsOverlap(ia, ib)) {
overlapping.add(a._id);
overlapping.add(b._id);
}
}
}
}
}
return overlapping;
}
function isValidTime(t: string): boolean {
return /^([01]\d|2[0-3]):([0-5]\d)$/.test(t);
}
function utcTimeNow(): string {
const now = new Date();
return now.toUTCString().slice(17, 22);
}
// Browsers may return "HH:MM:SS" from time inputs — strip seconds.
function normalizeTime(v: string): string {
return v.slice(0, 5);
}
let _nextId = 1;
function makeRow(partial: Partial<FeeTimeRange> = {}): RangeRow {
return {
_id: _nextId++,
start_time: partial.start_time ?? '',
end_time: partial.end_time ?? '',
provider_fee: partial.provider_fee ?? 1.05,
};
}
// ---------------------------------------------------------------------------
// Sub-component: read-only range list under a provider
// ---------------------------------------------------------------------------
function ProviderRangePreview({ schedules }: { schedules: FeeTimeRange[] }) {
if (schedules.length === 0) {
return (
<p className='text-muted-foreground pl-7 text-xs'>
No scheduled ranges default fee always applies.
</p>
);
}
return (
<ul className='space-y-0.5 pl-7'>
{schedules.map((s, i) => (
<li key={i} className='flex items-center gap-2 text-xs'>
<span className='text-muted-foreground font-mono'>
{s.start_time} {s.end_time} UTC
</span>
<Badge variant='outline' className='py-0 font-mono text-xs'>
×{s.provider_fee.toFixed(3)}
</Badge>
</li>
))}
</ul>
);
}
// ---------------------------------------------------------------------------
// Modal
// ---------------------------------------------------------------------------
export function ProviderFeeScheduleModal({
providers,
initialSelectedIds,
isOpen,
onClose,
onSuccess,
}: ProviderFeeScheduleModalProps) {
const queryClient = useQueryClient();
const [selectedIds, setSelectedIds] = useState<Set<number>>(new Set());
const [rows, setRows] = useState<RangeRow[]>([]);
const [enforceOverride, setEnforceOverride] = useState(false);
const [saving, setSaving] = useState(false);
const [clearing, setClearing] = useState(false);
// Reset state when modal opens. If a single provider is pre-selected,
// pre-populate the editor with its existing schedule so the user can edit it.
useEffect(() => {
if (!isOpen) return;
const ids =
initialSelectedIds && initialSelectedIds.length > 0
? new Set(initialSelectedIds)
: new Set(providers.map((p) => p.id));
setSelectedIds(ids);
setEnforceOverride(false);
if (initialSelectedIds && initialSelectedIds.length === 1) {
const provider = providers.find((p) => p.id === initialSelectedIds[0]);
const existing = provider?.provider_fee_schedules ?? [];
setRows(existing.length > 0 ? existing.map((s) => makeRow(s)) : []);
setEnforceOverride(true);
} else {
setRows([]);
}
// eslint-disable-next-line react-hooks/exhaustive-deps
}, [isOpen]);
const overlapping = findOverlappingIds(rows);
const allSelected =
providers.length > 0 && selectedIds.size === providers.length;
const noneSelected = selectedIds.size === 0;
const toggleProvider = (id: number) => {
setSelectedIds((prev) => {
const next = new Set(prev);
next.has(id) ? next.delete(id) : next.add(id);
return next;
});
};
const toggleAll = () => {
setSelectedIds(
allSelected ? new Set() : new Set(providers.map((p) => p.id))
);
};
const addRow = () => setRows((prev) => [...prev, makeRow()]);
const removeRow = (id: number) =>
setRows((prev) => prev.filter((r) => r._id !== id));
const updateRow = (
id: number,
field: keyof FeeTimeRange,
value: string | number
) =>
setRows((prev) =>
prev.map((r) => (r._id === id ? { ...r, [field]: value } : r))
);
const hasValidationErrors =
noneSelected ||
rows.some(
(r) =>
!isValidTime(r.start_time) ||
!isValidTime(r.end_time) ||
r.provider_fee <= 0
) ||
overlapping.size > 0;
const handleSave = async () => {
if (hasValidationErrors) return;
const newSchedules: FeeTimeRange[] = rows.map(
({ start_time, end_time, provider_fee }) => ({
start_time,
end_time,
provider_fee,
})
);
setSaving(true);
try {
await Promise.all(
[...selectedIds].map((id) => {
const provider = providers.find((p) => p.id === id);
const existing = provider?.provider_fee_schedules ?? [];
let finalSchedules: FeeTimeRange[];
if (enforceOverride) {
finalSchedules = newSchedules;
} else {
// Only override ranges that overlap with ANY of the new ranges.
// Keep existing non-overlapping ranges.
const keptExisting = existing.filter((ex) => {
return !newSchedules.some((nw) =>
_rangeIntervals(ex.start_time, ex.end_time).some((ia) =>
_rangeIntervals(nw.start_time, nw.end_time).some((ib) =>
_intervalsOverlap(ia, ib)
)
)
);
});
finalSchedules = [...keptExisting, ...newSchedules];
}
return AdminService.updateFeeSchedules(id, finalSchedules);
})
);
queryClient.invalidateQueries({ queryKey: ['upstream-providers'] });
toast.success(
`Fee schedules saved for ${selectedIds.size} provider${selectedIds.size > 1 ? 's' : ''}`
);
onSuccess();
onClose();
} catch (err) {
toast.error(
`Failed to save: ${err instanceof Error ? err.message : 'Unknown error'}`
);
} finally {
setSaving(false);
}
};
const handleClearAll = async () => {
if (noneSelected) return;
setClearing(true);
try {
await Promise.all(
[...selectedIds].map((id) => AdminService.deleteFeeSchedules(id))
);
queryClient.invalidateQueries({ queryKey: ['upstream-providers'] });
setRows([]);
toast.success(
`Fee schedules cleared for ${selectedIds.size} provider${selectedIds.size > 1 ? 's' : ''}`
);
onSuccess();
} catch (err) {
toast.error(
`Failed to clear: ${err instanceof Error ? err.message : 'Unknown error'}`
);
} finally {
setClearing(false);
}
};
return (
<Dialog open={isOpen} onOpenChange={(open) => !open && onClose()}>
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[660px]'>
<DialogHeader>
<DialogTitle>Fee Schedules</DialogTitle>
<DialogDescription>
Select providers and configure time-based fee ranges (UTC). Outside
scheduled ranges each provider&apos;s default fee applies. Current
UTC time:{' '}
<Badge variant='outline' className='font-mono'>
{utcTimeNow()}
</Badge>
</DialogDescription>
</DialogHeader>
{/* Provider selection with existing-range read view */}
<div className='space-y-2'>
<div className='flex items-center justify-between'>
<Label className='text-sm font-medium'>Apply to providers</Label>
<button
onClick={toggleAll}
className='text-muted-foreground hover:text-foreground text-xs underline-offset-2 hover:underline'
>
{allSelected ? 'Deselect all' : 'Select all'}
</button>
</div>
<div className='divide-y rounded-md border'>
{providers.map((p) => (
<div key={p.id} className='space-y-1.5 px-3 py-2'>
<label className='hover:bg-muted/50 flex cursor-pointer items-center gap-3 rounded'>
<Checkbox
checked={selectedIds.has(p.id)}
onCheckedChange={() => toggleProvider(p.id)}
/>
<span className='flex-1 text-sm font-medium'>
{p.provider_type}
</span>
<span className='text-muted-foreground truncate text-xs'>
{p.base_url}
</span>
</label>
<ProviderRangePreview
schedules={p.provider_fee_schedules ?? []}
/>
</div>
))}
</div>
{noneSelected && (
<p className='text-destructive text-xs'>
Select at least one provider.
</p>
)}
</div>
{/* Fee range editor */}
<div className='space-y-4'>
<div className='flex items-center justify-between'>
<Label className='text-sm font-medium'>
New schedule{' '}
{!enforceOverride && (
<span className='text-muted-foreground font-normal'>
(merges with existing ranges, overriding only overlaps)
</span>
)}
</Label>
<div className='flex items-center space-x-2'>
<Checkbox
id='enforce-override'
checked={enforceOverride}
onCheckedChange={(checked) => setEnforceOverride(!!checked)}
/>
<label
htmlFor='enforce-override'
className='text-xs leading-none font-medium peer-disabled:cursor-not-allowed peer-disabled:opacity-70'
>
Enforce overriding everything
</label>
</div>
</div>
<div className='space-y-2'>
{rows.length === 0 && (
<p className='text-muted-foreground rounded-md border border-dashed p-4 text-center text-sm'>
No ranges configured saving with no ranges will clear
schedules.
</p>
)}
{rows.map((row) => {
const isOverlap = overlapping.has(row._id);
const badTime =
(row.start_time && !isValidTime(row.start_time)) ||
(row.end_time && !isValidTime(row.end_time));
const badFee = row.provider_fee <= 1.0;
const hasError = isOverlap || badTime || badFee;
return (
<div
key={row._id}
className={`flex flex-col gap-2 rounded-md border p-3 sm:flex-row sm:items-end ${
hasError ? 'border-destructive bg-destructive/5' : ''
}`}
>
<div className='flex flex-1 flex-col gap-1'>
<Label className='text-xs'>Start (UTC)</Label>
<input
type='time'
value={row.start_time}
onChange={(e) =>
updateRow(
row._id,
'start_time',
normalizeTime(e.target.value)
)
}
className='border-input bg-background ring-offset-background focus-visible:ring-ring flex h-10 w-full rounded-md border px-3 py-2 font-mono text-sm focus-visible:ring-2 focus-visible:ring-offset-2 focus-visible:outline-none disabled:cursor-not-allowed disabled:opacity-50'
/>
</div>
<div className='flex flex-1 flex-col gap-1'>
<Label className='text-xs'>End (UTC)</Label>
<input
type='time'
value={row.end_time}
onChange={(e) =>
updateRow(
row._id,
'end_time',
normalizeTime(e.target.value)
)
}
className='border-input bg-background ring-offset-background focus-visible:ring-ring flex h-10 w-full rounded-md border px-3 py-2 font-mono text-sm focus-visible:ring-2 focus-visible:ring-offset-2 focus-visible:outline-none disabled:cursor-not-allowed disabled:opacity-50'
/>
</div>
<div className='flex flex-1 flex-col gap-1'>
<Label className='text-xs'>Fee multiplier</Label>
<Input
type='number'
step='0.001'
min='0.001'
placeholder='1.05'
value={row.provider_fee}
onChange={(e) =>
updateRow(
row._id,
'provider_fee',
parseFloat(e.target.value) || 0
)
}
/>
</div>
<Button
variant='ghost'
size='icon'
className='text-destructive hover:text-destructive shrink-0'
onClick={() => removeRow(row._id)}
>
<Trash2 className='h-4 w-4' />
</Button>
</div>
);
})}
</div>
{overlapping.size > 0 && (
<p className='text-destructive text-xs'>
Some ranges overlap fix them before saving.
</p>
)}
</div>
<div>
<Button variant='outline' size='sm' onClick={addRow}>
<Plus className='mr-1.5 h-4 w-4' />
Add Range
</Button>
</div>
<DialogFooter className='gap-2'>
<Button
variant='ghost'
onClick={handleClearAll}
disabled={clearing || noneSelected}
className='text-destructive hover:text-destructive mr-auto'
>
{clearing ? 'Clearing…' : 'Clear Selected'}
</Button>
<Button variant='outline' onClick={onClose}>
Cancel
</Button>
<Button onClick={handleSave} disabled={saving || hasValidationErrors}>
{saving
? 'Saving…'
: `Save to ${selectedIds.size} provider${selectedIds.size !== 1 ? 's' : ''}`}
</Button>
</DialogFooter>
</DialogContent>
</Dialog>
);
}

View File

@@ -202,26 +202,34 @@ export function ProviderFormFields({
<div className='grid gap-2'>
<Label htmlFor={`${idPrefix}provider_fee`}>
Provider Fee (Multiplier)
{mode === 'edit'
? 'Default Provider Fee (Multiplier)'
: 'Provider Fee (Multiplier)'}
</Label>
<Input
id={`${idPrefix}provider_fee`}
type='number'
step='0.001'
min='1.0'
value={formData.provider_fee || ''}
onChange={(e) =>
setFormData((prev) => ({
...prev,
provider_fee: e.target.value
? parseFloat(e.target.value)
: undefined,
}))
value={
(mode === 'edit'
? formData.provider_fee_default
: formData.provider_fee) || ''
}
onChange={(e) => {
const val = e.target.value ? parseFloat(e.target.value) : undefined;
setFormData((prev) =>
mode === 'edit'
? { ...prev, provider_fee_default: val }
: { ...prev, provider_fee: val }
);
}}
placeholder={providerFeePlaceholder}
/>
<p className='text-muted-foreground text-xs'>
1.01 means +1% e.g. currency exchange, card fees, etc.
{mode === 'edit'
? 'This is the default fee when no schedule is active. Updates will not affect currently active scheduled fees.'
: '1.01 means +1% e.g. currency exchange, card fees, etc.'}
</p>
</div>

View File

@@ -1,6 +1,6 @@
'use client';
import { useCallback, useState } from 'react';
import { useState } from 'react';
import Image from 'next/image';
import { Copy, Loader2, Zap, KeyRound } from 'lucide-react';
import { toast } from 'sonner';
@@ -10,7 +10,6 @@ import { Input } from '@/components/ui/input';
import { Textarea } from '@/components/ui/textarea';
import { Label } from '@/components/ui/label';
import { Badge } from '@/components/ui/badge';
import { Separator } from '@/components/ui/separator';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
interface RoutstrCreateKeySectionProps {

View File

@@ -0,0 +1,265 @@
'use client';
import * as React from 'react';
import { useState, useEffect, useCallback } from 'react';
import {
AdminService,
type CliTokenListItem,
type CliTokenCreated,
} from '@/lib/api/services/admin';
import {
Card,
CardContent,
CardHeader,
CardTitle,
CardDescription,
} from '@/components/ui/card';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Label } from '@/components/ui/label';
import { Skeleton } from '@/components/ui/skeleton';
import { Alert, AlertDescription } from '@/components/ui/alert';
import { AlertCircle, Copy, Trash2, Check } from 'lucide-react';
import { toast } from 'sonner';
function formatTs(ts: number | null): string {
if (!ts) return '—';
return new Date(ts * 1000).toLocaleString();
}
export function CliTokensSettings(): React.ReactElement {
const [tokens, setTokens] = useState<CliTokenListItem[]>([]);
const [loading, setLoading] = useState(true);
const [error, setError] = useState<string | null>(null);
const [name, setName] = useState('');
const [expiresInDays, setExpiresInDays] = useState<string>('');
const [creating, setCreating] = useState(false);
const [newToken, setNewToken] = useState<CliTokenCreated | null>(null);
const [copied, setCopied] = useState(false);
const loadTokens = useCallback(async (): Promise<void> => {
setLoading(true);
setError(null);
try {
const data = await AdminService.listCliTokens();
setTokens(data);
} catch (err: unknown) {
const message =
err instanceof Error ? err.message : 'Failed to load tokens';
setError(message);
} finally {
setLoading(false);
}
}, []);
useEffect(() => {
void loadTokens();
}, [loadTokens]);
async function handleCreate(): Promise<void> {
const trimmed = name.trim();
if (!trimmed) {
toast.error('Name is required');
return;
}
const days = expiresInDays.trim()
? Number.parseInt(expiresInDays.trim(), 10)
: undefined;
if (days !== undefined && (Number.isNaN(days) || days <= 0)) {
toast.error('Expiry must be a positive number of days');
return;
}
setCreating(true);
try {
const created = await AdminService.createCliToken(trimmed, days);
setNewToken(created);
setName('');
setExpiresInDays('');
await loadTokens();
toast.success('Token created. Copy it now — it will not be shown again.');
} catch (err: unknown) {
const message =
err instanceof Error ? err.message : 'Failed to create token';
toast.error(message);
} finally {
setCreating(false);
}
}
async function handleRevoke(id: string): Promise<void> {
if (
!confirm('Revoke this token? Any CLI/agent using it will lose access.')
) {
return;
}
try {
await AdminService.revokeCliToken(id);
await loadTokens();
toast.success('Token revoked');
} catch (err: unknown) {
const message =
err instanceof Error ? err.message : 'Failed to revoke token';
toast.error(message);
}
}
async function handleCopy(): Promise<void> {
if (!newToken) return;
await navigator.clipboard.writeText(newToken.token);
setCopied(true);
setTimeout(() => setCopied(false), 2000);
}
return (
<div className='space-y-6'>
<Card>
<CardHeader>
<CardTitle>Create CLI Token</CardTitle>
<CardDescription>
Generate a long-lived bearer token for the Routstr CLI or AI agents.
Use this token in <code>~/.routstr/config.json</code> or with{' '}
<code>routstr init --token &lt;token&gt;</code>.
</CardDescription>
</CardHeader>
<CardContent className='space-y-4'>
{newToken && (
<Alert className='border-green-500/50 bg-green-500/10'>
<AlertDescription className='space-y-3'>
<div className='font-medium text-green-700 dark:text-green-400'>
Token created. Copy it now it will not be shown again.
</div>
<div className='flex items-center gap-2'>
<code className='bg-muted flex-1 rounded px-3 py-2 text-xs break-all'>
{newToken.token}
</code>
<Button
type='button'
variant='outline'
size='sm'
onClick={handleCopy}
>
{copied ? (
<Check className='h-4 w-4' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</div>
<Button
type='button'
variant='ghost'
size='sm'
onClick={() => setNewToken(null)}
>
Dismiss
</Button>
</AlertDescription>
</Alert>
)}
<div className='grid grid-cols-1 gap-4 md:grid-cols-2'>
<div className='space-y-2'>
<Label htmlFor='cli-token-name'>Name</Label>
<Input
id='cli-token-name'
placeholder='e.g. dev-laptop, ci-runner'
value={name}
onChange={(e) => setName(e.target.value)}
disabled={creating}
/>
</div>
<div className='space-y-2'>
<Label htmlFor='cli-token-expiry'>
Expires in days (optional)
</Label>
<Input
id='cli-token-expiry'
type='number'
min='1'
placeholder='Never expires if blank'
value={expiresInDays}
onChange={(e) => setExpiresInDays(e.target.value)}
disabled={creating}
/>
</div>
</div>
<Button onClick={handleCreate} disabled={creating || !name.trim()}>
{creating ? 'Creating…' : 'Create Token'}
</Button>
</CardContent>
</Card>
<Card>
<CardHeader>
<CardTitle>Active Tokens</CardTitle>
<CardDescription>
Tokens authorize CLI/agent calls to admin endpoints. Revoke any
token that may have been exposed.
</CardDescription>
</CardHeader>
<CardContent>
{error && (
<Alert variant='destructive' className='mb-4'>
<AlertCircle className='h-4 w-4' />
<AlertDescription>{error}</AlertDescription>
</Alert>
)}
{loading ? (
<div className='space-y-2'>
<Skeleton className='h-12 w-full' />
<Skeleton className='h-12 w-full' />
</div>
) : tokens.length === 0 ? (
<p className='text-muted-foreground text-sm'>
No tokens yet. Create one above.
</p>
) : (
<div className='overflow-x-auto'>
<table className='w-full text-sm'>
<thead>
<tr className='text-muted-foreground border-b text-left'>
<th className='py-2 pr-4 font-medium'>Name</th>
<th className='py-2 pr-4 font-medium'>Token</th>
<th className='py-2 pr-4 font-medium'>Created</th>
<th className='py-2 pr-4 font-medium'>Last used</th>
<th className='py-2 pr-4 font-medium'>Expires</th>
<th className='py-2 font-medium'></th>
</tr>
</thead>
<tbody>
{tokens.map((t) => (
<tr key={t.id} className='border-b last:border-0'>
<td className='py-2 pr-4'>{t.name}</td>
<td className='py-2 pr-4 font-mono text-xs'>
{t.token_preview}
</td>
<td className='text-muted-foreground py-2 pr-4'>
{formatTs(t.created_at)}
</td>
<td className='text-muted-foreground py-2 pr-4'>
{formatTs(t.last_used_at)}
</td>
<td className='text-muted-foreground py-2 pr-4'>
{t.expires_at ? formatTs(t.expires_at) : 'Never'}
</td>
<td className='py-2'>
<Button
type='button'
variant='ghost'
size='sm'
onClick={() => void handleRevoke(t.id)}
>
<Trash2 className='h-4 w-4' />
</Button>
</td>
</tr>
))}
</tbody>
</table>
</div>
)}
</CardContent>
</Card>
</div>
);
}

View File

@@ -12,6 +12,12 @@ export const ProviderTypeSchema = z.object({
can_show_balance: z.boolean(),
});
export const FeeTimeRangeSchema = z.object({
start_time: z.string(),
end_time: z.string(),
provider_fee: z.number(),
});
export const UpstreamProviderSchema = z.object({
id: z.number(),
provider_type: z.string(),
@@ -20,7 +26,9 @@ export const UpstreamProviderSchema = z.object({
api_version: z.string().nullable().optional(),
enabled: z.boolean(),
provider_fee: z.number().optional(),
provider_fee_default: z.number().optional(),
provider_settings: z.record(z.string(), z.any()).nullable().optional(),
provider_fee_schedules: z.array(FeeTimeRangeSchema).optional().default([]),
});
export const CreateUpstreamProviderSchema = z.object({
@@ -30,6 +38,7 @@ export const CreateUpstreamProviderSchema = z.object({
api_version: z.string().nullable().optional(),
enabled: z.boolean().default(true),
provider_fee: z.number().optional(),
provider_fee_default: z.number().optional(),
provider_settings: z.record(z.string(), z.any()).nullable().optional(),
});
@@ -40,6 +49,7 @@ export const UpdateUpstreamProviderSchema = z.object({
api_version: z.string().nullable().optional(),
enabled: z.boolean().optional(),
provider_fee: z.number().optional(),
provider_fee_default: z.number().optional(),
provider_settings: z.record(z.string(), z.any()).nullable().optional(),
});
@@ -76,6 +86,7 @@ export const AdminModelSchema = z.object({
canonical_slug: z.string().nullable().optional(),
alias_ids: z.array(z.string()).nullable().optional(),
enabled: z.boolean().default(true),
forwarded_model_id: z.string().nullable().optional(),
});
export const ProviderModelsSchema = z.object({
@@ -96,6 +107,7 @@ export type CreateUpstreamProvider = z.infer<
export type UpdateUpstreamProvider = z.infer<
typeof UpdateUpstreamProviderSchema
>;
export type FeeTimeRange = z.infer<typeof FeeTimeRangeSchema>;
export type AdminModel = z.infer<typeof AdminModelSchema>;
export type AdminModelPricing = z.infer<typeof AdminModelPricingSchema>;
export type AdminModelArchitecture = z.infer<
@@ -307,6 +319,30 @@ export class AdminService {
);
}
static async getFeeSchedules(providerId: number): Promise<FeeTimeRange[]> {
return await apiClient.get<FeeTimeRange[]>(
`/admin/api/upstream-providers/${providerId}/fee-schedules`
);
}
static async updateFeeSchedules(
providerId: number,
schedules: FeeTimeRange[]
): Promise<FeeTimeRange[]> {
return await apiClient.put<FeeTimeRange[]>(
`/admin/api/upstream-providers/${providerId}/fee-schedules`,
{ schedules }
);
}
static async deleteFeeSchedules(
providerId: number
): Promise<{ ok: boolean }> {
return await apiClient.delete<{ ok: boolean }>(
`/admin/api/upstream-providers/${providerId}/fee-schedules`
);
}
static async getProviderModels(providerId: number): Promise<ProviderModels> {
const data = await apiClient.get<ProviderModels>(
`/admin/api/upstream-providers/${providerId}/models`
@@ -890,13 +926,17 @@ export class AdminService {
type?: string,
status?: string,
search?: string,
limit: number = 100
source?: string,
limit: number = 50,
offset: number = 0
): Promise<TransactionsResponse> {
const params = new URLSearchParams();
if (type) params.append('type', type);
if (status) params.append('status', status);
if (search) params.append('search', search);
if (source) params.append('source', source);
params.append('limit', limit.toString());
params.append('offset', offset.toString());
return await apiClient.get<TransactionsResponse>(
`/admin/api/transactions?${params.toString()}`
@@ -961,6 +1001,45 @@ export class AdminService {
balance_data: number | null | Record<string, unknown>;
}>(`/admin/api/upstream-providers/${providerId}/balance`);
}
// ── CLI Tokens ──
static async listCliTokens(): Promise<CliTokenListItem[]> {
return await apiClient.get<CliTokenListItem[]>('/admin/api/cli-tokens');
}
static async createCliToken(
name: string,
expiresInDays?: number
): Promise<CliTokenCreated> {
return await apiClient.post<CliTokenCreated>('/admin/api/cli-tokens', {
name,
expires_in_days: expiresInDays ?? null,
});
}
static async revokeCliToken(tokenId: string): Promise<{ ok: boolean }> {
return await apiClient.delete<{ ok: boolean }>(
`/admin/api/cli-tokens/${encodeURIComponent(tokenId)}`
);
}
}
export interface CliTokenListItem {
id: string;
name: string;
token_preview: string;
created_at: number;
last_used_at: number | null;
expires_at: number | null;
}
export interface CliTokenCreated {
id: string;
name: string;
token: string;
created_at: number;
expires_at: number | null;
}
export const TemporaryBalanceSchema = z.object({
@@ -1134,6 +1213,8 @@ export interface Transaction {
created_at: number;
collected: boolean;
swept: boolean;
source: 'x-cashu' | 'apikey';
api_key_hashed_key?: string;
}
export interface TransactionsResponse {

2864
uv.lock generated

File diff suppressed because it is too large Load Diff