Compare 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
9qeklajcandGitHub 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
9qeklajcandGitHub 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
9qeklajcandGitHub 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
9qeklajcandGitHub 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
9qeklajcandGitHub 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
9qeklajcandGitHub 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
9qeklajcandGitHub 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
9qeklajcandGitHub 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
9qeklajcandGitHub 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
root 4691294b24 Merge branch 'main' into dynamic-prices 2026-04-20 13:10:00 +02:00
9qeklajc d84d249f2c Merge branch 'main' into add-dev-cut
# Conflicts:
#	routstr/auth.py
2026-04-15 23:07:37 +02:00
9qeklajc 111060c004 Merge branch 'add-cost-usage-to-messages-endpoint' into dynamic-prices 2026-04-15 21:32:38 +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 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 85532f32a1 Merge branch 'main' into dynamic-prices 2026-04-14 00:46:00 +02:00
9qeklajc ea8257c44b Merge branch 'main' into dynamic-prices 2026-04-14 00:32:22 +02:00
9qeklajc d89e740ff8 Merge branch 'main' into dynamic-prices
# Conflicts:
#	routstr/upstream/base.py
2026-04-13 23:52:46 +02:00
9qeklajc af4be5bfec Merge branch 'main' into dynamic-prices 2026-04-12 22:46:18 +02:00
9qeklajc df2a925577 added-dynamic-price-setting 2026-04-12 17:38:35 +02:00
9qeklajc 453337cb2c update default payout 2026-04-05 00:36:59 +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
33 changed files with 3158 additions and 330 deletions
@@ -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")
@@ -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")
+34
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")
@@ -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
+1 -1
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"
+22 -1
View File
@@ -12,7 +12,7 @@ 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,
@@ -25,6 +25,12 @@ 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)
# PREPAID_BALANCE = int(os.environ.get("PREPAID_BALANCE", "0")) * 1000 # Convert to msats
@@ -741,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(
@@ -832,6 +849,7 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(cost.total_msats)
payments_logger.info(
"FINALIZE",
extra={
@@ -935,6 +953,7 @@ async def adjust_payment_for_tokens(
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
@@ -1005,6 +1024,7 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
@@ -1130,6 +1150,7 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
+193 -57
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
@@ -556,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
@@ -566,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)])
@@ -615,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,
@@ -624,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(
@@ -645,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(
@@ -682,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)
@@ -690,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(
@@ -720,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."""
+68 -2
View File
@@ -12,7 +12,7 @@ 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
@@ -220,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:
+13 -4
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)
+113
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
+42
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()
@@ -118,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:
@@ -133,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
@@ -140,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
+232 -134
View File
@@ -67,6 +67,7 @@ class BaseUpstreamProvider:
api_key: str
provider_fee: float = 1.05
_models_cache: list[Model] = []
_raw_models_cache: list[Model] = []
_models_by_id: dict[str, Model] = {}
def __init__(self, base_url: str, api_key: str, provider_fee: float = 1.01):
@@ -81,6 +82,7 @@ class BaseUpstreamProvider:
self.api_key = api_key
self.provider_fee = provider_fee
self._models_cache = []
self._raw_models_cache = []
self._models_by_id = {}
@classmethod
@@ -604,57 +606,83 @@ class BaseUpstreamProvider:
)
yield prefix + part
# Stream finished, process usage if found
if usage_chunk_data:
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
try:
cost_data = await adjust_payment_for_tokens(
fresh_key,
usage_chunk_data,
session,
max_cost_for_model,
)
remaining_balance_msats = fresh_key.balance
# Merge cost into usage
usage_chunk_data["usage"]["cost"] = cost_data.get(
"total_usd", 0.0
)
usage_chunk_data["usage"]["cost_sats"] = (
cost_data.get("total_msats", 0) // 1000
)
usage_chunk_data["usage"]["remaining_balance_msats"] = (
remaining_balance_msats
)
# Keep detailed cost in metadata
usage_chunk_data["metadata"] = usage_chunk_data.get(
"metadata", {}
)
usage_chunk_data["metadata"]["routstr"] = {
"cost": cost_data
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
cost_data: dict
try:
adjustment_input = (
usage_chunk_data
if usage_chunk_data is not None
else {
"model": last_model_seen or "unknown",
"usage": None,
}
usage_chunk_data["metadata"]["routstr"]["cost"][
"sats_cost"
] = cost_data.get("total_msats", 0) // 1000
usage_chunk_data["metadata"]["routstr"]["cost"][
"remaining_balance_msats"
] = remaining_balance_msats
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
usage_finalized = True
except Exception as e:
logger.exception(
"Error during usage finalization",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
)
# Fallback: yield original usage chunk if adjustment fails
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
)
cost_data = await adjust_payment_for_tokens(
fresh_key,
adjustment_input,
session,
max_cost_for_model,
)
usage_finalized = True
except Exception as e:
logger.exception(
"Error during usage finalization",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
)
if not usage_finalized:
await finalize_db_only()
# Fall back so we still emit a non-zero sats cost downstream.
cost_data = {
"base_msats": 0,
"input_msats": 0,
"output_msats": 0,
"total_msats": 0,
"total_usd": 0.0,
"input_tokens": 0,
"output_tokens": 0,
}
if usage_chunk_data is None:
if not hasattr(self, "_current_stream_id"):
self._current_stream_id = (
f"chatcmpl-{uuid.uuid4()}"
)
usage_chunk_data = {
"id": self._current_stream_id,
"object": "chat.completion.chunk",
"model": last_model_seen or "unknown",
"choices": [],
"usage": {
"prompt_tokens": cost_data.get(
"input_tokens", 0
),
"completion_tokens": cost_data.get(
"output_tokens", 0
),
"total_tokens": cost_data.get(
"input_tokens", 0
)
+ cost_data.get("output_tokens", 0),
},
}
try:
self.inject_cost_metadata(
usage_chunk_data, cost_data, fresh_key
)
except Exception:
logger.exception(
"Failed to inject cost metadata into streaming chunk",
extra={
"key_hash": key.hashed_key[:8] + "...",
},
)
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
if done_seen:
yield b"data: [DONE]\n\n"
@@ -926,65 +954,108 @@ class BaseUpstreamProvider:
)
yield prefix + part
# Stream finished, process usage if found
if usage_chunk_data:
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
try:
cost_data = await adjust_payment_for_tokens(
fresh_key,
usage_chunk_data,
session,
max_cost_for_model,
)
remaining_balance_msats = fresh_key.balance
# Merge cost into usage chunk
if (
"response" in usage_chunk_data
and "usage" in usage_chunk_data["response"]
):
usage_chunk_data["response"]["usage"]["cost"] = (
cost_data.get("total_usd", 0.0)
)
usage_chunk_data["response"]["usage"][
"cost_sats"
] = cost_data.get("total_msats", 0) // 1000
usage_chunk_data["response"]["usage"][
"remaining_balance_msats"
] = remaining_balance_msats
elif "usage" in usage_chunk_data:
usage_chunk_data["usage"]["cost"] = cost_data.get(
"total_usd", 0.0
)
usage_chunk_data["usage"]["cost_sats"] = (
cost_data.get("total_msats", 0) // 1000
)
usage_chunk_data["usage"][
"remaining_balance_msats"
] = remaining_balance_msats
# Keep detailed cost in metadata
usage_chunk_data["metadata"] = usage_chunk_data.get(
"metadata", {}
)
usage_chunk_data["metadata"]["routstr"] = {
"cost": cost_data
# Always emit a cost-bearing data chunk
async with create_session() as session:
fresh_key = await session.get(key.__class__, key.hashed_key)
if fresh_key:
cost_data: dict
try:
adjustment_input = (
usage_chunk_data
if usage_chunk_data is not None
else {
"model": last_model_seen or "unknown",
"usage": None,
}
usage_chunk_data["metadata"]["routstr"]["cost"][
"sats_cost"
] = cost_data.get("total_msats", 0) // 1000
usage_chunk_data["metadata"]["routstr"]["cost"][
"remaining_balance_msats"
] = remaining_balance_msats
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
usage_finalized = True
except Exception:
# Fallback: yield original usage chunk if adjustment fails
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
)
cost_data = await adjust_payment_for_tokens(
fresh_key,
adjustment_input,
session,
max_cost_for_model,
)
usage_finalized = True
except Exception as e:
logger.exception(
"Error during Responses API usage finalization",
extra={
"key_hash": key.hashed_key[:8] + "...",
"error": str(e),
},
)
cost_data = {
"base_msats": 0,
"input_msats": 0,
"output_msats": 0,
"total_msats": 0,
"total_usd": 0.0,
"input_tokens": 0,
"output_tokens": 0,
}
if not usage_finalized:
await finalize_db_only()
if usage_chunk_data is None:
usage_chunk_data = {
"type": "response.completed",
"response": {
"model": last_model_seen or "unknown",
"usage": {
"input_tokens": cost_data.get(
"input_tokens", 0
),
"output_tokens": cost_data.get(
"output_tokens", 0
),
"total_tokens": cost_data.get(
"input_tokens", 0
)
+ cost_data.get("output_tokens", 0),
},
},
"usage": {
"input_tokens": cost_data.get(
"input_tokens", 0
),
"output_tokens": cost_data.get(
"output_tokens", 0
),
"total_tokens": cost_data.get(
"input_tokens", 0
)
+ cost_data.get("output_tokens", 0),
},
}
remaining_balance_msats = fresh_key.balance
sats_cost = cost_data.get("total_msats", 0) // 1000
if (
"response" in usage_chunk_data
and isinstance(usage_chunk_data["response"], dict)
and "usage" in usage_chunk_data["response"]
):
usage_chunk_data["response"]["usage"]["cost"] = (
cost_data.get("total_usd", 0.0)
)
usage_chunk_data["response"]["usage"][
"cost_sats"
] = sats_cost
usage_chunk_data["response"]["usage"][
"remaining_balance_msats"
] = remaining_balance_msats
try:
self.inject_cost_metadata(
usage_chunk_data, cost_data, fresh_key
)
except Exception:
logger.exception(
"Failed to inject cost metadata into Responses streaming chunk",
extra={
"key_hash": key.hashed_key[:8] + "...",
},
)
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
if done_seen:
yield b"data: [DONE]\n\n"
@@ -1308,7 +1379,8 @@ class BaseUpstreamProvider:
)
usage_finalized = True
yield f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
# Emit the full combined_data as the cost
yield f"event: cost\ndata: {json.dumps(combined_data)}\n\n".encode()
except Exception:
pass
@@ -3712,8 +3784,28 @@ class BaseUpstreamProvider:
None,
)
async def refresh_models_cache(self) -> None:
"""Refresh the in-memory models cache from upstream API."""
def apply_fee_to_cache(self) -> None:
"""Apply current provider_fee to raw models and update active cache."""
models_with_fees = [
self._apply_provider_fee_to_model(m) for m in self._raw_models_cache
]
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}
async def refresh_models_cache(self, skip_network: bool = False) -> None:
"""Refresh the in-memory models cache from upstream API and database.
Args:
skip_network: If True, only refresh from database, skip hitting upstream API.
"""
try:
async with create_session() as session:
stmt = select(UpstreamProviderRow).where(
@@ -3727,6 +3819,9 @@ class BaseUpstreamProvider:
if not provider or not provider.id:
raise HTTPException(status_code=404, detail="Provider not found")
# Update fee from DB if it changed
self.provider_fee = provider.provider_fee
db_models = await list_models(
session=session,
upstream_id=provider.id,
@@ -3734,34 +3829,37 @@ class BaseUpstreamProvider:
apply_fees=False,
)
db_model_ids: set[str] = {model.id for model in db_models}
models = await self.fetch_models()
model_ids = [model.id for model in models]
diff = set(db_model_ids) - set(model_ids)
for db_model_id in diff:
found_db_model = next(
(
model_obj
for model_obj in db_models
if model_obj.id == db_model_id
if skip_network:
# Use existing raw models but filter/merge with DB models
# This avoids hitting the network
current_raw = {m.id: m for m in self._raw_models_cache}
# Keep only those still in current_raw (if we wanted to be strict)
# but actually we want to merge with db_models
models = []
# Add all db_models (they take precedence as overrides)
models.extend(db_models)
# Add current raw models that are not in DB
for m_id, m in current_raw.items():
if m_id not in db_model_ids:
models.append(m)
else:
models = await self.fetch_models()
model_ids = [model.id for model in models]
diff = set(db_model_ids) - set(model_ids)
for db_model_id in diff:
found_db_model = next(
(
model_obj
for model_obj in db_models
if model_obj.id == db_model_id
)
)
)
models.append(found_db_model)
models.append(found_db_model)
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}
self._raw_models_cache = models
self.apply_fee_to_cache()
except Exception as e:
logger.error(
+13 -4
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:
+1 -103
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,
)
+40
View File
@@ -589,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]
+302
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)
+190
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
@@ -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]:
@@ -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
+64
View File
@@ -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(
+256
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
+117
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]
+9 -1
View File
@@ -51,7 +51,15 @@ async def test_stream_with_id_injection() -> None:
base.adjust_payment_for_tokens = AsyncMock(
return_value={"total_usd": 0.1, "total_msats": 100}
)
base.create_session = MagicMock()
# 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,
+43 -7
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>
);
}
+5
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>
-1
View File
@@ -339,7 +339,6 @@ export default function TransactionsPage() {
setApikeyPage(0);
}, [type, status, search]);
const activeQuery = activeTab === 'x-cashu' ? xcashuQuery : apikeyQuery;
const isRefetching = xcashuQuery.isRefetching || apikeyQuery.isRefetching;
const renderCardContent = (
+3 -3
View File
@@ -540,7 +540,7 @@ export function AddProviderModelDialog({
};
return (
<FormItem>
<FormLabel>Upstream Model ID</FormLabel>
<FormLabel>Client Alias ID</FormLabel>
<FormControl>
<div className='flex gap-2'>
<Input
@@ -565,8 +565,8 @@ export function AddProviderModelDialog({
</div>
</FormControl>
<FormDescription>
Model ID sent to the upstream provider. Defaults to the
model&apos;s own ID.
Alternate ID that clients can use to reference this
model. Defaults to the model&apos;s own ID.
</FormDescription>
<FormMessage />
</FormItem>
+27
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'
@@ -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>
);
}
+18 -10
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>
@@ -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>
);
}
+74
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(),
});
@@ -97,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<
@@ -308,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`
@@ -966,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({
Generated
+1 -1
View File
@@ -1878,7 +1878,7 @@ wheels = [
[[package]]
name = "routstr"
version = "0.4.1"
version = "0.4.3"
source = { editable = "." }
dependencies = [
{ name = "aiosqlite" },