Compare commits

..
29 changed files with 2042 additions and 513 deletions
-11
View File
@@ -14,14 +14,6 @@ jobs:
steps:
- name: Checkout code
uses: actions/checkout@v3
with:
fetch-depth: 0
- name: Resolve git metadata
id: gitmeta
run: |
echo "sha=$(git rev-parse --short=7 HEAD)" >> "$GITHUB_OUTPUT"
echo "tag=$(git describe --tags --exact-match HEAD 2>/dev/null || true)" >> "$GITHUB_OUTPUT"
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v2
@@ -38,9 +30,6 @@ jobs:
with:
context: .
push: true
build-args: |
GIT_COMMIT=${{ steps.gitmeta.outputs.sha }}
GIT_TAG=${{ steps.gitmeta.outputs.tag }}
tags: |
ghcr.io/routstr/proxy:latest
ghcr.io/routstr/core:latest
-4
View File
@@ -21,10 +21,6 @@ WORKDIR /app
COPY . .
ARG GIT_COMMIT=""
ARG GIT_TAG=""
ENV GIT_COMMIT=${GIT_COMMIT}
ENV GIT_TAG=${GIT_TAG}
ENV PORT=8000
ENV PYTHONUNBUFFERED=1
-4
View File
@@ -41,10 +41,6 @@ RUN uv sync --no-dev
# Copy the built UI from the ui-builder stage
COPY --from=ui-builder /app/ui/out ./ui_out
ARG GIT_COMMIT=""
ARG GIT_TAG=""
ENV GIT_COMMIT=${GIT_COMMIT}
ENV GIT_TAG=${GIT_TAG}
ENV PORT=8000
ENV PYTHONUNBUFFERED=1
@@ -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")
+2 -2
View File
@@ -317,13 +317,13 @@ async def validate_bearer_key(
extra={"key_hash": hashed_key[:8] + "..."},
)
logger.debug(
logger.info(
"AUTH: About to call credit_balance",
extra={"token_preview": bearer_key[:50]},
)
try:
msats = await credit_balance(bearer_key, new_key, session)
logger.debug(
logger.info(
"AUTH: credit_balance returned successfully", extra={"msats": msats}
)
except Exception as credit_error:
+105 -52
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,
@@ -639,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
@@ -649,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)])
@@ -698,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,
@@ -707,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(
@@ -728,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(
@@ -765,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)
@@ -773,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(
@@ -803,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."""
+9 -2
View File
@@ -79,10 +79,11 @@ class ApiKey(SQLModel, table=True): # type: ignore
async def reset_all_reserved_balances(session: AsyncSession) -> None:
logger.info("Resetting all reserved balances to 0")
stmt = update(ApiKey).values(reserved_balance=0)
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
logger.info("Reset reserved balances on startup")
logger.info("Reserved balances reset successfully")
class ModelRow(SQLModel, table=True): # type: ignore
@@ -219,11 +220,17 @@ 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"},
+9 -13
View File
@@ -22,20 +22,16 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon
# Get status code and detail - works for both FastAPI and Starlette HTTPException
status_code = getattr(exc, "status_code", 500)
detail = getattr(exc, "detail", str(exc))
path = request.url.path
# 4xx is client behaviour; the uvicorn access log already records it.
# Only 5xx warrants a server-side warning/error log here.
if status_code >= 500:
logger.error(
f"HTTP {status_code} on {path}: {detail}",
extra={
"request_id": request_id,
"status_code": status_code,
"detail": detail,
"path": path,
},
)
logger.warning(
"HTTP exception",
extra={
"request_id": request_id,
"status_code": status_code,
"detail": detail,
"path": request.url.path,
},
)
return JSONResponse(
status_code=status_code,
+9 -34
View File
@@ -41,23 +41,14 @@ import logging.config
import logging.handlers
import os
import re
import sys
import tomllib
from datetime import datetime
from pathlib import Path
from typing import Any
from pythonjsonlogger import jsonlogger
from rich.console import Console
from rich.logging import RichHandler
# Only use RichHandler when stdout is a real TTY. In non-TTY contexts
# (docker logs, pipes, CI) Rich pads every line to width and wraps long
# records, producing visually-empty trailing whitespace and split records.
# A plain StreamHandler avoids both problems.
_stdout_is_tty = sys.stdout.isatty()
_console = Console(soft_wrap=True) if _stdout_is_tty else None
# Define custom TRACE level
TRACE_LEVEL = 5
logging.addLevelName(TRACE_LEVEL, "TRACE")
@@ -270,26 +261,6 @@ def setup_logging() -> None:
if console_enabled:
handlers.append("console")
if _stdout_is_tty:
console_handler: dict[str, Any] = {
"()": RichHandler,
"level": log_level,
"show_time": False,
"show_path": False,
"rich_tracebacks": True,
"markup": True,
"console": _console,
"filters": ["request_id_filter", "security_filter"],
}
else:
console_handler = {
"class": "logging.StreamHandler",
"level": log_level,
"formatter": "plain",
"stream": "ext://sys.stdout",
"filters": ["request_id_filter", "security_filter"],
}
LOGGING_CONFIG = {
"version": 1,
"disable_existing_loggers": False,
@@ -299,10 +270,6 @@ def setup_logging() -> None:
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s",
"datefmt": "%Y-%m-%d %H:%M:%S",
},
"plain": {
"format": "%(asctime)s %(levelname)-7s %(name)s %(message)s",
"datefmt": "%Y-%m-%d %H:%M:%S",
},
},
"filters": {
"version_filter": {"()": VersionFilter},
@@ -310,7 +277,15 @@ def setup_logging() -> None:
"security_filter": {"()": SecurityFilter},
},
"handlers": {
"console": console_handler,
"console": {
"()": RichHandler,
"level": log_level,
"show_time": False,
"show_path": False,
"rich_tracebacks": True,
"markup": True,
"filters": ["request_id_filter", "security_filter"],
},
"file": {
"()": DailyRotatingFileHandler,
"level": log_level,
+98 -79
View File
@@ -1,4 +1,5 @@
import asyncio
import os
from contextlib import asynccontextmanager
from pathlib import Path
from typing import AsyncGenerator
@@ -8,8 +9,6 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, RedirectResponse
from fastapi.staticfiles import StaticFiles
from starlette.exceptions import HTTPException
from starlette.responses import Response as StarletteResponse
from starlette.types import Scope
from ..auth import periodic_key_reset
from ..balance import balance_router, deprecated_wallet_router
@@ -31,12 +30,16 @@ from .logging import get_logger, setup_logging
from .middleware import LoggingMiddleware
from .settings import SettingsService
from .settings import settings as global_settings
from .version import __version__
# Initialize logging first
setup_logging()
logger = get_logger(__name__)
if os.getenv("VERSION_SUFFIX") is not None:
__version__ = f"0.4.3-{os.getenv('VERSION_SUFFIX')}"
else:
__version__ = "0.4.3"
@asynccontextmanager
async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
@@ -194,23 +197,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
)
class _ImmutableStaticFiles(StaticFiles):
"""Static files with long Cache-Control for content-hashed Next.js assets.
Files under `/_next/static/` are emitted with content hashes in their
filenames and never mutate, so we serve them with a one-year immutable
cache header so browsers and CDNs stop revalidating on every reload.
"""
async def get_response(self, path: str, scope: Scope) -> StarletteResponse:
response = await super().get_response(path, scope)
if response.status_code == 200:
response.headers["Cache-Control"] = (
"public, max-age=31536000, immutable"
)
return response
app = FastAPI(version=__version__, lifespan=lifespan)
@@ -257,7 +243,7 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
app.mount(
"/_next",
_ImmutableStaticFiles(directory=UI_DIST_PATH / "_next", check_dir=True),
StaticFiles(directory=UI_DIST_PATH / "_next", check_dir=True),
name="next-static",
)
@@ -265,70 +251,100 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
async def serve_root_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "index.html")
# Serve the App Router RSC payload for the home page.
# Add explicit route for /index.txt to redirect to /
@app.get("/index.txt", include_in_schema=False)
async def serve_root_rsc() -> FileResponse:
return FileResponse(
UI_DIST_PATH / "index.txt", media_type="text/x-component"
)
# Next.js is built with `trailingSlash: true`, so all UI page URLs end
# with a slash (e.g. `/login/`). The proxy router catches `/{path:path}`
# before FastAPI's `redirect_slashes` logic can normalize the URL, so we
# must register both the with-slash and without-slash variants here.
UI_PAGES = (
"dashboard",
"login",
"model",
"providers",
"settings",
"transactions",
"balances",
"logs",
"usage",
"unauthorized",
)
def _register_ui_page(name: str) -> None:
page_dir = UI_DIST_PATH / name
index_html = page_dir / "index.html"
index_txt = page_dir / "index.txt"
async def serve_page() -> FileResponse:
return FileResponse(index_html)
async def serve_page_rsc() -> FileResponse:
return FileResponse(index_txt, media_type="text/x-component")
app.add_api_route(
f"/{name}",
serve_page,
methods=["GET"],
include_in_schema=False,
name=f"serve_{name}_ui",
)
app.add_api_route(
f"/{name}/",
serve_page,
methods=["GET"],
include_in_schema=False,
name=f"serve_{name}_ui_slash",
)
app.add_api_route(
f"/{name}/index.txt",
serve_page_rsc,
methods=["GET"],
include_in_schema=False,
name=f"serve_{name}_rsc",
)
for _page in UI_PAGES:
_register_ui_page(_page)
async def redirect_index_txt() -> RedirectResponse:
return RedirectResponse("/")
@app.get("/admin")
async def admin_redirect() -> FileResponse:
return FileResponse(UI_DIST_PATH / "index.html")
@app.get("/dashboard", include_in_schema=False)
async def serve_dashboard_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "index.html")
@app.get("/login", include_in_schema=False)
async def serve_login_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "login" / "index.html")
# Add explicit route for /login/index.txt to redirect to /login
@app.get("/login/index.txt", include_in_schema=False)
async def redirect_login_index_txt() -> RedirectResponse:
return RedirectResponse("/login")
@app.get("/model", include_in_schema=False)
async def serve_models_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "model" / "index.html")
# Add explicit route for /model/index.txt to redirect to /model
@app.get("/model/index.txt", include_in_schema=False)
async def redirect_model_index_txt() -> RedirectResponse:
return RedirectResponse("/model")
@app.get("/providers", include_in_schema=False)
async def serve_providers_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "providers" / "index.html")
# Add explicit route for /providers/index.txt to redirect to /providers
@app.get("/providers/index.txt", include_in_schema=False)
async def redirect_providers_index_txt() -> RedirectResponse:
return RedirectResponse("/providers")
@app.get("/settings", include_in_schema=False)
async def serve_settings_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "settings" / "index.html")
# Add explicit route for /settings/index.txt to redirect to /settings
@app.get("/settings/index.txt", include_in_schema=False)
async def redirect_settings_index_txt() -> RedirectResponse:
return RedirectResponse("/settings")
@app.get("/transactions", include_in_schema=False)
async def serve_transactions_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "transactions" / "index.html")
# Add explicit route for /transactions/index.txt to redirect to /transactions
@app.get("/transactions/index.txt", include_in_schema=False)
async def redirect_transactions_index_txt() -> RedirectResponse:
return RedirectResponse("/transactions")
@app.get("/balances", include_in_schema=False)
async def serve_balances_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "balances" / "index.html")
# Add explicit route for /balances/index.txt to redirect to /balances
@app.get("/balances/index.txt", include_in_schema=False)
async def redirect_balances_index_txt() -> RedirectResponse:
return RedirectResponse("/balances")
@app.get("/logs", include_in_schema=False)
async def serve_logs_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "logs" / "index.html")
# Add explicit route for /logs/index.txt to redirect to /logs
@app.get("/logs/index.txt", include_in_schema=False)
async def redirect_logs_index_txt() -> RedirectResponse:
return RedirectResponse("/logs")
@app.get("/usage", include_in_schema=False)
async def serve_usage_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "usage" / "index.html")
# Add explicit route for /usage/index.txt to redirect to /usage
@app.get("/usage/index.txt", include_in_schema=False)
async def redirect_usage_index_txt() -> RedirectResponse:
return RedirectResponse("/usage")
@app.get("/unauthorized", include_in_schema=False)
async def serve_unauthorized_ui() -> FileResponse:
return FileResponse(UI_DIST_PATH / "unauthorized" / "index.html")
# Add explicit route for /unauthorized/index.txt to redirect to /unauthorized
@app.get("/unauthorized/index.txt", include_in_schema=False)
async def redirect_unauthorized_index_txt() -> RedirectResponse:
return RedirectResponse("/unauthorized")
@app.get("/favicon.ico", include_in_schema=False)
async def serve_favicon() -> FileResponse:
icon_path = UI_DIST_PATH / "icon.ico"
@@ -340,6 +356,9 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
async def serve_icon() -> FileResponse:
return FileResponse(UI_DIST_PATH / "icon.ico")
app.mount(
"/static", StaticFiles(directory=UI_DIST_PATH, check_dir=True), name="ui-static"
)
else:
logger.warning(
f"UI dist directory not found at {UI_DIST_PATH}, skipping static file serving"
+63 -69
View File
@@ -14,54 +14,8 @@ logger = get_logger(__name__)
request_id_context: ContextVar[str | None] = ContextVar("request_id")
# Methods that are never logged: HEAD requests are health probes from
# monitoring/load balancers, OPTIONS are CORS preflights — both are framework
# chatter, not user-meaningful events.
_SKIP_LOG_METHODS: frozenset[str] = frozenset({"HEAD", "OPTIONS"})
# Path prefixes to skip. Includes Next.js static chunks and the admin
# dashboard's internal polling API (/admin/api/*) which the UI hits on a timer
# to refresh balances, logs, providers, etc. — high volume, low diagnostic
# value. Mutating admin actions are recorded separately in the audit log.
_SKIP_LOG_PREFIXES: tuple[str, ...] = (
"/_next/",
"/admin/api/",
)
# Exact paths to skip. RSC payload prefetches (`*/index.txt`) fire automatically
# as the user hovers near `<Link>`s, and `/v1/wallet/info` is polled by the UI.
_SKIP_LOG_EXACT: frozenset[str] = frozenset(
{
"/favicon.ico",
"/icon.ico",
"/v1/wallet/info",
"/index.txt",
"/login/index.txt",
"/model/index.txt",
"/providers/index.txt",
"/settings/index.txt",
"/transactions/index.txt",
"/balances/index.txt",
"/logs/index.txt",
"/usage/index.txt",
"/unauthorized/index.txt",
}
)
def _should_log(method: str, path: str) -> bool:
if method in _SKIP_LOG_METHODS:
return False
if path in _SKIP_LOG_EXACT:
return False
return not any(path.startswith(prefix) for prefix in _SKIP_LOG_PREFIXES)
class LoggingMiddleware(BaseHTTPMiddleware):
"""Middleware to log proxy interactions and page navigation.
Skips logging for static assets and Next.js chunks to avoid noise.
"""
"""Middleware to log detailed request and response information."""
async def dispatch(self, request: Request, call_next: Callable) -> Response:
# Generate request ID
@@ -71,20 +25,56 @@ class LoggingMiddleware(BaseHTTPMiddleware):
# Set request ID in context for logging
token = request_id_context.set(request_id)
path = request.url.path
should_log = _should_log(request.method, path)
# Start timing
start_time = time.time()
if should_log:
logger.info(
"Incoming request",
# Log request details
request_body = None
if request.method in ["POST", "PUT", "PATCH"]:
try:
# Only read body for non-streaming requests
if hasattr(request, "_body"):
request_body = await request.body()
except Exception:
pass
# Log incoming request
logger.info(
"Incoming request",
extra={
"request_id": request_id,
"method": request.method,
"path": request.url.path,
"query_params": dict(request.query_params),
"headers": {
k: v
for k, v in request.headers.items()
if k.lower()
not in [
"authorization",
"x-cashu",
"cookie",
"cf-connecting-ip",
"cf-ipcountry",
"x-forwarded-for",
"x-real-ip",
]
},
"body_size": len(request_body) if request_body else 0,
},
)
# Log at TRACE level for full body (security filter will redact sensitive data)
if request_body and hasattr(logger, "exception"):
logger.exception(
"Request body",
extra={
"request_id": request_id,
"method": request.method,
"path": path,
"query_params": dict(request.query_params),
"path": request.url.path,
"body": request_body.decode("utf-8", errors="ignore")[
:1000
], # Limit size
},
)
@@ -92,32 +82,36 @@ class LoggingMiddleware(BaseHTTPMiddleware):
try:
response = await call_next(request)
if should_log:
duration = time.time() - start_time
logger.info(
"Request completed",
extra={
"request_id": request_id,
"method": request.method,
"path": path,
"status_code": response.status_code,
"duration_ms": round(duration * 1000, 2),
},
)
# Calculate duration
duration = time.time() - start_time
# Log response
logger.info(
"Request completed",
extra={
"request_id": request_id,
"method": request.method,
"path": request.url.path,
"status_code": response.status_code,
"duration_ms": round(duration * 1000, 2),
},
)
if hasattr(response, "headers"):
response.headers["x-routstr-request-id"] = request_id
return response
except Exception as e:
# Always log failures, even for skipped paths, so we don't lose errors.
# Calculate duration
duration = time.time() - start_time
# Log error
logger.error(
"Request failed",
extra={
"request_id": request_id,
"method": request.method,
"path": path,
"path": request.url.path,
"duration_ms": round(duration * 1000, 2),
"error": str(e),
"error_type": type(e).__name__,
-81
View File
@@ -1,81 +0,0 @@
"""Application version resolution.
Priority order:
1. ``VERSION_SUFFIX`` env var (manual override; preserves prior behaviour).
2. Bare base version when HEAD is on the matching release tag (detected via
``GIT_TAG`` env or ``git describe --tags --exact-match HEAD``).
3. ``GIT_COMMIT`` env var (build-time injection) -> ``<base>+g<sha>``.
4. Local ``.git`` lookup (source checkouts) -> ``<base>+g<sha>``.
5. Fallback: bare base version.
The ``+g<sha>`` form is PEP 440 local-version syntax so the result remains a
valid package version.
"""
from __future__ import annotations
import os
import subprocess
from functools import lru_cache
from pathlib import Path
BASE_VERSION = "0.4.3"
_REPO_ROOT = Path(__file__).resolve().parents[2]
_GIT_TIMEOUT_SECONDS = 2.0
def _run_git(*args: str) -> str | None:
try:
result = subprocess.run( # noqa: S603 - fixed argv, no shell
["git", *args],
cwd=_REPO_ROOT,
check=False,
capture_output=True,
text=True,
timeout=_GIT_TIMEOUT_SECONDS,
)
except (FileNotFoundError, subprocess.SubprocessError, OSError):
return None
if result.returncode != 0:
return None
return result.stdout.strip() or None
def _git_short_sha() -> str | None:
sha = os.getenv("GIT_COMMIT", "").strip()
if sha:
return sha[:7]
return _run_git("rev-parse", "--short=7", "HEAD")
def _on_tagged_release() -> bool:
tag = os.getenv("GIT_TAG", "").strip()
if tag:
return tag.lstrip("v") == BASE_VERSION
described = _run_git("describe", "--tags", "--exact-match", "HEAD")
if not described:
return False
return described.lstrip("v") == BASE_VERSION
@lru_cache(maxsize=1)
def get_version() -> str:
suffix = os.getenv("VERSION_SUFFIX")
if suffix is not None:
return f"{BASE_VERSION}-{suffix}"
if _on_tagged_release():
return BASE_VERSION
sha = _git_short_sha()
if not sha:
return BASE_VERSION
return f"{BASE_VERSION}+g{sha}"
__version__ = get_version()
__all__ = ["BASE_VERSION", "__version__", "get_version"]
+1 -1
View File
@@ -26,7 +26,7 @@ logger = get_logger(__name__)
def get_app_version() -> str | None:
try:
from ..core.version import __version__ as imported_version
from ..core.main import __version__ as imported_version
return imported_version
except Exception:
+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
+1 -7
View File
@@ -350,9 +350,6 @@ async def _update_sats_pricing_once() -> None:
from ..proxy import get_upstreams, refresh_model_maps
upstreams = get_upstreams()
if not upstreams:
return
sats_to_usd = sats_usd_price()
updated_count = 0
@@ -366,10 +363,7 @@ async def _update_sats_pricing_once() -> None:
updated_count += len(updated_models)
if updated_count > 0:
logger.info(
f"Updated sats pricing for {updated_count} models",
extra={"models_updated": updated_count},
)
logger.info("Updated sats pricing", extra={"models_updated": updated_count})
await refresh_model_maps()
+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
+55 -27
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
@@ -3782,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(
@@ -3797,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,
@@ -3804,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(
+3 -4
View File
@@ -197,14 +197,13 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
existing_providers = result.all()
if not existing_providers:
logger.info(
"No upstream providers found in database, seeding from settings"
)
await _seed_providers_from_settings(session, settings)
await session.commit()
result = await session.exec(select(UpstreamProviderRow))
existing_providers = result.all()
if existing_providers:
logger.info(
f"Seeded {len(existing_providers)} upstream providers from settings"
)
async def _init_single_provider(
provider_row: UpstreamProviderRow,
+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,
)
+2 -2
View File
@@ -326,7 +326,7 @@ async def credit_balance(
except Exception:
pass
logger.debug(
logger.info(
"Cashu token successfully redeemed and stored",
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
)
@@ -488,7 +488,7 @@ async def fetch_all_balances(
async def periodic_payout() -> None:
if not settings.receive_ln_address:
logger.warning("RECEIVE_LN_ADDRESS is not set, periodic payout disabled")
logger.error("RECEIVE_LN_ADDRESS is not set, skipping payout")
return
while True:
await asyncio.sleep(60 * 15)
+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
+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
+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>
);
}
+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>
+35
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`