diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index aca29594..ae676584 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -62,14 +62,18 @@ jobs: - name: Setup Node.js uses: actions/setup-node@v4 with: - node-version: '18' - cache: 'npm' + node-version: "18" + cache: "npm" cache-dependency-path: ui/package-lock.json - name: Install UI dependencies working-directory: ./ui run: npm ci + - name: Run UI format check + working-directory: ./ui + run: npm run format-check + - name: Run UI linting working-directory: ./ui run: npm run lint diff --git a/migrations/versions/lightning_invoices_table.py b/migrations/versions/lightning_invoices_table.py new file mode 100644 index 00000000..0b795e8d --- /dev/null +++ b/migrations/versions/lightning_invoices_table.py @@ -0,0 +1,39 @@ +"""Add lightning_invoices table + +Revision ID: lightning_invoices +Revises: a1a1a1a1a1a1 +Create Date: 2025-12-10 21:00:00.000000 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +revision = "lightning_invoices" +down_revision = "a1a1a1a1a1a1" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "lightning_invoices", + sa.Column("id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("bolt11", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("amount_sats", sa.Integer(), nullable=False), + sa.Column("description", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("payment_hash", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("status", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("api_key_hash", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("purpose", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("created_at", sa.Integer(), nullable=False), + sa.Column("expires_at", sa.Integer(), nullable=False), + sa.Column("paid_at", sa.Integer(), nullable=True), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("bolt11"), + sa.UniqueConstraint("payment_hash"), + ) + + +def downgrade() -> None: + op.drop_table("lightning_invoices") diff --git a/pyproject.toml b/pyproject.toml index 34eb16e6..a25cbbaf 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,6 +20,7 @@ dependencies = [ "nostr>=0.0.2", "mdurl==0.1.2", "pillow>=10", + "openai>=1.98.0", ] [dependency-groups] diff --git a/routstr/algorithm.py b/routstr/algorithm.py index ccdc7088..515c2576 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -150,18 +150,6 @@ def should_prefer_model( # Prefer lower adjusted cost should_replace = candidate_adjusted < current_adjusted - # Log provider changes when candidate wins - if should_replace: - candidate_provider_name = getattr( - candidate_provider, "provider_type", "unknown" - ) - current_provider_name = getattr(current_provider, "provider_type", "unknown") - logger.debug( - f"Model selection for alias '{alias}': choosing {candidate_provider_name} " - f"(cost: ${candidate_adjusted:.6f}) over {current_provider_name} " - f"(cost: ${current_adjusted:.6f})" - ) - return should_replace @@ -254,7 +242,12 @@ def create_model_mappings( # Add to unique models base_id = get_base_model_id(model_to_use.id) if not is_openrouter or base_id not in unique_models: - unique_model = model_to_use.copy(update={"id": base_id}) + unique_model = model_to_use.copy( + update={ + "id": base_id, + "upstream_provider_id": upstream.provider_type, + } + ) unique_models[base_id] = unique_model # Get all aliases for this model @@ -289,12 +282,8 @@ def create_model_mappings( provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1 logger.debug( - "Created model mappings", - extra={ - "unique_model_count": len(unique_models), - "total_alias_count": len(model_instances), - "provider_distribution": provider_counts, - }, + f"Updated model mappings with ({len(unique_models)} unique models and {len(model_instances)} aliases)", + extra={"provider_distribution": provider_counts}, ) return model_instances, provider_map, unique_models diff --git a/routstr/balance.py b/routstr/balance.py index 883c4673..62358761 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -10,6 +10,7 @@ from .auth import validate_bearer_key from .core.db import ApiKey, AsyncSession, get_session from .core.logging import get_logger from .core.settings import settings +from .lightning import lightning_router from .wallet import credit_balance, recieve_token, send_to_lnurl, send_token router = APIRouter() @@ -238,6 +239,8 @@ async def wallet_catch_all(path: str) -> NoReturn: ) +balance_router.include_router(lightning_router) balance_router.include_router(router) + deprecated_wallet_router = APIRouter(prefix="/v1/wallet", include_in_schema=False) deprecated_wallet_router.include_router(router) diff --git a/routstr/core/db.py b/routstr/core/db.py index c6effbe3..f40212b3 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -1,4 +1,5 @@ import os +import time from contextlib import asynccontextmanager from typing import AsyncGenerator @@ -71,6 +72,28 @@ class ModelRow(SQLModel, table=True): # type: ignore upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models") +class LightningInvoice(SQLModel, table=True): # type: ignore + __tablename__ = "lightning_invoices" + + id: str = Field(primary_key=True, description="Unique invoice identifier") + bolt11: str = Field(description="BOLT11 invoice string", unique=True) + amount_sats: int = Field(description="Amount in satoshis") + description: str = Field(description="Invoice description") + payment_hash: str = Field(description="Payment hash for tracking", unique=True) + status: str = Field( + default="pending", description="pending, paid, expired, cancelled" + ) + api_key_hash: str | None = Field( + default=None, description="Associated API key hash for topup operations" + ) + purpose: str = Field(description="create or topup") + created_at: int = Field( + default_factory=lambda: int(time.time()), description="Unix timestamp" + ) + expires_at: int = Field(description="Unix timestamp when invoice expires") + paid_at: int | None = Field(default=None, description="Unix timestamp when paid") + + class UpstreamProviderRow(SQLModel, table=True): # type: ignore __tablename__ = "upstream_providers" id: int | None = Field(default=None, primary_key=True) @@ -126,8 +149,6 @@ def run_migrations() -> None: import pathlib try: - logger.info("Starting database migrations") - # Get the path to the alembic.ini file project_root = pathlib.Path(__file__).resolve().parents[2] alembic_ini_path = project_root / "alembic.ini" @@ -144,7 +165,6 @@ def run_migrations() -> None: alembic_cfg.set_main_option("sqlalchemy.url", DATABASE_URL) # Run migrations to the latest revision - logger.info("Running migrations to latest revision") command.upgrade(alembic_cfg, "head") logger.info("Database migrations completed successfully") diff --git a/routstr/core/logging.py b/routstr/core/logging.py index 65b6a5c0..00474949 100644 --- a/routstr/core/logging.py +++ b/routstr/core/logging.py @@ -338,6 +338,11 @@ def setup_logging() -> None: "handlers": ["console"] if console_enabled else [], "propagate": False, }, + "openai": { + "level": "WARNING", + "handlers": ["console"] if console_enabled else [], + "propagate": False, + }, "httpcore": { "level": "WARNING", "handlers": ["console"] if console_enabled else [], @@ -360,6 +365,11 @@ def setup_logging() -> None: }, "watchfiles.main": {"level": "WARNING", "handlers": [], "propagate": False}, "aiosqlite": {"level": "ERROR", "handlers": [], "propagate": False}, + "alembic": { + "level": "WARNING", + "handlers": ["console"] if console_enabled else [], + "propagate": False, + }, }, "root": { "level": log_level, diff --git a/routstr/core/main.py b/routstr/core/main.py index 06b86d65..d2974d4a 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -54,9 +54,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: try: # Run database migrations on startup - # This ensures the database schema is always up-to-date in production - # Migrations are idempotent - running them multiple times is safe - logger.info("Running database migrations") run_migrations() # Initialize database connection pools @@ -104,6 +101,9 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: yield + except asyncio.CancelledError: + # Expected during shutdown + pass except Exception as e: logger.error( "Application startup failed", diff --git a/routstr/discovery.py b/routstr/discovery.py index a2680145..03a82703 100644 --- a/routstr/discovery.py +++ b/routstr/discovery.py @@ -62,7 +62,9 @@ async def query_nostr_relay_for_providers( if data[0] == "EVENT" and data[1] == sub_id: event = data[2] - logger.debug(f"Found provider announcement: {event['id']}") + logger.debug( + f"Found provider announcement: {event['id'][:6]}...{event['id'][-6:]}" + ) events.append(event) elif data[0] == "EOSE" and data[1] == sub_id: logger.debug("Received EOSE message") diff --git a/routstr/lightning.py b/routstr/lightning.py new file mode 100644 index 00000000..50d1608d --- /dev/null +++ b/routstr/lightning.py @@ -0,0 +1,270 @@ +import hashlib +import secrets +import time + +from fastapi import APIRouter, Depends, HTTPException +from pydantic import BaseModel, Field +from sqlmodel import select +from sqlmodel.ext.asyncio.session import AsyncSession + +from .core.db import ApiKey, LightningInvoice, get_session +from .core.logging import get_logger +from .core.settings import settings +from .wallet import get_wallet + +logger = get_logger(__name__) + +lightning_router = APIRouter(prefix="/lightning") + + +class InvoiceCreateRequest(BaseModel): + amount_sats: int = Field(gt=0, le=1_000_000, description="Amount in satoshis") + purpose: str = Field(description="create or topup", pattern="^(create|topup)$") + api_key: str | None = Field( + default=None, description="Required for topup operations" + ) + + +class InvoiceCreateResponse(BaseModel): + invoice_id: str + bolt11: str + amount_sats: int + expires_at: int + payment_hash: str + + +class InvoiceStatusResponse(BaseModel): + status: str + api_key: str | None = None + amount_sats: int + paid_at: int | None = None + created_at: int + expires_at: int + + +class InvoiceRecoverRequest(BaseModel): + bolt11: str = Field(description="BOLT11 invoice string") + + +async def generate_lightning_invoice( + amount_sats: int, description: str +) -> tuple[str, str]: + wallet = await get_wallet(settings.primary_mint, "sat") + quote = await wallet.request_mint(amount_sats) + return quote.request, quote.quote + + +def generate_invoice_id() -> str: + return secrets.token_urlsafe(16) + + +@lightning_router.post("/invoice", response_model=InvoiceCreateResponse) +async def create_invoice( + request: InvoiceCreateRequest, + session: AsyncSession = Depends(get_session), +) -> InvoiceCreateResponse: + if request.purpose == "topup" and not request.api_key: + raise HTTPException( + status_code=400, detail="api_key is required for topup operations" + ) + + if request.purpose == "topup" and request.api_key: + if not request.api_key.startswith("sk-"): + raise HTTPException(status_code=400, detail="Invalid API key format") + + api_key = await session.get(ApiKey, request.api_key[3:]) + if not api_key: + raise HTTPException(status_code=404, detail="API key not found") + + try: + description = f"Routstr {request.purpose} {request.amount_sats} sats" + bolt11, payment_hash = await generate_lightning_invoice( + request.amount_sats, description + ) + + invoice_id = generate_invoice_id() + expires_at = int(time.time()) + 3600 # 1 hour expiry + + invoice = LightningInvoice( + id=invoice_id, + bolt11=bolt11, + amount_sats=request.amount_sats, + description=description, + payment_hash=payment_hash, + status="pending", + api_key_hash=request.api_key[3:] if request.api_key else None, + purpose=request.purpose, + expires_at=expires_at, + ) + + session.add(invoice) + await session.commit() + + logger.info( + "Lightning invoice created", + extra={ + "invoice_id": invoice_id, + "amount_sats": request.amount_sats, + "purpose": request.purpose, + "expires_at": expires_at, + }, + ) + + return InvoiceCreateResponse( + invoice_id=invoice_id, + bolt11=bolt11, + amount_sats=request.amount_sats, + expires_at=expires_at, + payment_hash=payment_hash, + ) + + except Exception as e: + logger.error(f"Failed to create Lightning invoice: {e}") + raise HTTPException( + status_code=500, detail="Failed to create Lightning invoice" + ) + + +@lightning_router.get( + "/invoice/{invoice_id}/status", response_model=InvoiceStatusResponse +) +async def get_invoice_status( + invoice_id: str, + session: AsyncSession = Depends(get_session), +) -> InvoiceStatusResponse: + invoice = await session.get(LightningInvoice, invoice_id) + if not invoice: + raise HTTPException(status_code=404, detail="Invoice not found") + + if invoice.status == "pending" and int(time.time()) > invoice.expires_at: + invoice.status = "expired" + await session.commit() + + if invoice.status == "pending": + await check_invoice_payment(invoice, session) + + api_key = None + if invoice.status == "paid" and invoice.purpose == "create": + if invoice.api_key_hash: + api_key = f"sk-{invoice.api_key_hash}" + elif ( + invoice.status == "paid" and invoice.purpose == "topup" and invoice.api_key_hash + ): + api_key = f"sk-{invoice.api_key_hash}" + + return InvoiceStatusResponse( + status=invoice.status, + api_key=api_key, + amount_sats=invoice.amount_sats, + paid_at=invoice.paid_at, + created_at=invoice.created_at, + expires_at=invoice.expires_at, + ) + + +@lightning_router.post("/recover", response_model=InvoiceStatusResponse) +async def recover_invoice( + request: InvoiceRecoverRequest, + session: AsyncSession = Depends(get_session), +) -> InvoiceStatusResponse: + result = await session.exec( + select(LightningInvoice).where(LightningInvoice.bolt11 == request.bolt11) + ) + invoice = result.first() + + if not invoice: + raise HTTPException(status_code=404, detail="Invoice not found") + + if invoice.status == "pending": + await check_invoice_payment(invoice, session) + + api_key = None + if invoice.status == "paid": + if invoice.purpose == "create" and invoice.api_key_hash: + api_key = f"sk-{invoice.api_key_hash}" + elif invoice.purpose == "topup" and invoice.api_key_hash: + api_key = f"sk-{invoice.api_key_hash}" + + return InvoiceStatusResponse( + status=invoice.status, + api_key=api_key, + amount_sats=invoice.amount_sats, + paid_at=invoice.paid_at, + created_at=invoice.created_at, + expires_at=invoice.expires_at, + ) + + +async def check_invoice_payment( + invoice: LightningInvoice, session: AsyncSession +) -> None: + try: + wallet = await get_wallet(settings.primary_mint, "sat") + + mint_status = await wallet.get_mint_quote(invoice.payment_hash) + + if mint_status.paid: + invoice.status = "paid" + invoice.paid_at = int(time.time()) + + if invoice.purpose == "create": + api_key = await create_api_key_from_invoice(invoice, session) + invoice.api_key_hash = api_key.hashed_key + elif invoice.purpose == "topup" and invoice.api_key_hash: + await topup_api_key_from_invoice(invoice, session) + + await session.commit() + + logger.info( + "Lightning invoice paid", + extra={ + "invoice_id": invoice.id, + "amount_sats": invoice.amount_sats, + "purpose": invoice.purpose, + "api_key_hash": invoice.api_key_hash[:8] + "..." + if invoice.api_key_hash + else None, + }, + ) + + except Exception as e: + logger.error(f"Failed to check invoice payment: {e}") + + +async def create_api_key_from_invoice( + invoice: LightningInvoice, session: AsyncSession +) -> ApiKey: + wallet = await get_wallet(settings.primary_mint, "sat") + await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash) + + dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}" + hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest() + + api_key = ApiKey( + hashed_key=hashed_key, + balance=invoice.amount_sats * 1000, # Convert to msats + refund_currency="sat", + refund_mint_url=settings.primary_mint, + ) + + session.add(api_key) + await session.flush() + + return api_key + + +async def topup_api_key_from_invoice( + invoice: LightningInvoice, session: AsyncSession +) -> None: + wallet = await get_wallet(settings.primary_mint, "sat") + await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash) + + if not invoice.api_key_hash: + raise ValueError("No API key associated with topup invoice") + + api_key = await session.get(ApiKey, invoice.api_key_hash) + if not api_key: + raise ValueError("Associated API key not found") + + api_key.balance += invoice.amount_sats * 1000 # Convert to msats + await session.flush() diff --git a/routstr/nip91.py b/routstr/nip91.py index 3a0a40d2..d14c2c6e 100644 --- a/routstr/nip91.py +++ b/routstr/nip91.py @@ -215,7 +215,7 @@ async def query_nip91_events( continue events_out.append(ev_dict) logger.debug( - f"Found existing NIP-91 event: {ev_dict.get('id', '')}" + f"Found listing event: {ev_dict.get('id', '')[:6]}...{ev_dict.get('id', '')[-6:]}" ) if drained: last_event_ts = time.time() diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 0403f007..0831ca37 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -1,7 +1,8 @@ import asyncio import json import random -from typing import Final +from pathlib import Path +from urllib.request import urlopen import httpx from fastapi import APIRouter, Depends @@ -18,14 +19,6 @@ logger = get_logger(__name__) models_router = APIRouter() -DEFAULT_EXCLUDED_MODEL_IDS: Final[set[str]] = { - "openrouter/auto", - "google/gemini-2.5-pro-exp-03-25", - "opengvlab/internvl3-78b", - "openrouter/sonoma-dusk-alpha", - "openrouter/sonoma-sky-alpha", -} - class Architecture(BaseModel): modality: str @@ -38,10 +31,12 @@ class Architecture(BaseModel): class Pricing(BaseModel): prompt: float completion: float - request: float - image: float - web_search: float - internal_reasoning: float + request: float = 0.0 + image: float = 0.0 + web_search: float = 0.0 + internal_reasoning: float = 0.0 + input_cache_read: float = 0.0 + input_cache_write: float = 0.0 max_prompt_cost: float = 0.0 # in sats not msats max_completion_cost: float = 0.0 # in sats not msats max_cost: float = 0.0 # in sats not msats @@ -65,7 +60,7 @@ class Model(BaseModel): per_request_limits: dict | None = None top_provider: TopProvider | None = None enabled: bool = True - upstream_provider_id: int | None = None + upstream_provider_id: int | str | None = None canonical_slug: str | None = None alias_ids: list[str] | None = None @@ -73,6 +68,27 @@ class Model(BaseModel): return hash(self.id) +def _has_valid_pricing(model: dict) -> bool: + """Check if model has valid pricing (not free, no negative values).""" + pricing = model.get("pricing", {}) + if not pricing: + return False + + try: + prompt = float(pricing.get("prompt", 0)) + completion = float(pricing.get("completion", 0)) + except (ValueError, TypeError): + return False + + if prompt < 0 or completion < 0: + return False + + if prompt == 0 and completion == 0: + return False + + return True + + async def async_fetch_openrouter_models(source_filter: str | None = None) -> list[dict]: """Asynchronously fetch model information from OpenRouter API.""" base_url = "https://openrouter.ai/api/v1" @@ -96,10 +112,10 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis model["id"] = model_id[len(source_prefix) :] model_id = model["id"] - if ( - "(free)" in model.get("name", "") - or model_id in DEFAULT_EXCLUDED_MODEL_IDS - ): + if "(free)" in model.get("name", ""): + continue + + if not _has_valid_pricing(model): continue models_data.append(model) diff --git a/routstr/payment/price.py b/routstr/payment/price.py index c20ac34e..9011ce1d 100644 --- a/routstr/payment/price.py +++ b/routstr/payment/price.py @@ -110,10 +110,6 @@ async def _update_prices() -> None: return BTC_USD_PRICE = btc_price SATS_USD_PRICE = btc_price / 100_000_000 - logger.info( - "Updated BTC/USD price", - extra={"btc_usd": btc_price, "sats_usd": SATS_USD_PRICE}, - ) def btc_usd_price() -> float: diff --git a/routstr/upstream/__init__.py b/routstr/upstream/__init__.py index 074cd14d..13c26791 100644 --- a/routstr/upstream/__init__.py +++ b/routstr/upstream/__init__.py @@ -2,6 +2,7 @@ from .anthropic import AnthropicUpstreamProvider from .azure import AzureUpstreamProvider from .base import BaseUpstreamProvider from .fireworks import FireworksUpstreamProvider +from .gemini import GeminiUpstreamProvider from .generic import GenericUpstreamProvider from .groq import GroqUpstreamProvider from .ollama import OllamaUpstreamProvider @@ -15,6 +16,7 @@ upstream_provider_classes: list[type[BaseUpstreamProvider]] = [ AnthropicUpstreamProvider, AzureUpstreamProvider, FireworksUpstreamProvider, + GeminiUpstreamProvider, GenericUpstreamProvider, GroqUpstreamProvider, OllamaUpstreamProvider, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index d891c7ac..7726b5f6 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1726,7 +1726,6 @@ class BaseUpstreamProvider: Returns: List of Model objects with pricing """ - logger.debug(f"Fetching models for {self.provider_type or self.base_url}") try: or_models, provider_models_response = await asyncio.gather( @@ -1753,18 +1752,9 @@ class BaseUpstreamProvider: else: not_found_models.append(model_id) - logger.info( - "Fetched models for provider", - extra={ - "provider": self.provider_type or self.base_url, - "found_count": len(found_models), - "not_found_count": len(not_found_models), - }, - ) - if not_found_models: logger.debug( - "Models not found in OpenRouter", + f"({len(not_found_models)}/{len(provider_model_ids)}) unmatched models for {self.provider_type or self.base_url}", extra={"not_found_models": not_found_models}, ) @@ -1832,10 +1822,7 @@ class BaseUpstreamProvider: 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.provider_type or self.base_url}", - extra={"model_count": len(models)}, - ) + except Exception as e: logger.error( f"Failed to refresh models cache for {self.provider_type or self.base_url}", @@ -1904,14 +1891,15 @@ class BaseUpstreamProvider: f"Provider {self.provider_type} does not support top-up" ) - async def get_balance(self) -> dict[str, object]: + async def get_balance(self) -> float | None: """Get the current account balance from the provider. Returns: - Dict with balance information + Float representing the balance amount, or None if not supported/available. + Typically in USD or the provider's credit unit. Raises: - NotImplementedError: If provider does not support balance checking + NotImplementedError: If provider does not support balance checking (default behavior) """ raise NotImplementedError( f"Provider {self.provider_type} does not support balance checking" diff --git a/routstr/upstream/clients/__init__.py b/routstr/upstream/clients/__init__.py new file mode 100644 index 00000000..2443feaa --- /dev/null +++ b/routstr/upstream/clients/__init__.py @@ -0,0 +1,3 @@ +from .gemini import GeminiClient + +__all__ = ["GeminiClient"] diff --git a/routstr/upstream/clients/base.py b/routstr/upstream/clients/base.py new file mode 100644 index 00000000..17fab428 --- /dev/null +++ b/routstr/upstream/clients/base.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +from abc import ABC, abstractmethod +from typing import Any, AsyncGenerator + + +class BaseAPIClient(ABC): + """Base class for AI provider API clients.""" + + def __init__(self, api_key: str, base_url: str | None = None): + self.api_key = api_key + self.base_url = base_url + + @abstractmethod + async def generate_content( + self, + model: str, + messages: list[dict[str, Any]], + temperature: float | None = None, + max_tokens: int | None = None, + **kwargs: Any, + ) -> dict[str, Any]: + """Generate content non-streaming.""" + pass + + @abstractmethod + def generate_content_stream( + self, + model: str, + messages: list[dict[str, Any]], + temperature: float | None = None, + max_tokens: int | None = None, + **kwargs: Any, + ) -> AsyncGenerator[dict[str, Any], None]: + pass + + @abstractmethod + async def list_models(self) -> list[dict[str, Any]]: + """List available models.""" + pass diff --git a/routstr/upstream/clients/gemini.py b/routstr/upstream/clients/gemini.py new file mode 100644 index 00000000..893dd647 --- /dev/null +++ b/routstr/upstream/clients/gemini.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +from typing import Any, AsyncGenerator + +from openai import AsyncOpenAI + +from .base import BaseAPIClient + + +class GeminiClient(BaseAPIClient): + """Gemini API client using OpenAI compatibility layer.""" + + def __init__(self, api_key: str, base_url: str | None = None): + super().__init__(api_key, base_url) + self.client = AsyncOpenAI( + api_key=api_key, + base_url=base_url + or "https://generativelanguage.googleapis.com/v1beta/openai/", + ) + + async def generate_content( + self, + model: str, + messages: list[dict[str, Any]], + temperature: float | None = None, + max_tokens: int | None = None, + **kwargs: Any, + ) -> dict[str, Any]: + from openai import NOT_GIVEN + + response = await self.client.chat.completions.create( + model=model, + messages=messages, # type: ignore + temperature=temperature if temperature is not None else NOT_GIVEN, + max_tokens=max_tokens if max_tokens is not None else NOT_GIVEN, + top_p=kwargs.get("top_p", NOT_GIVEN), + ) + return response.model_dump() + + async def generate_content_stream( + self, + model: str, + messages: list[dict[str, Any]], + temperature: float | None = None, + max_tokens: int | None = None, + **kwargs: Any, + ) -> AsyncGenerator[dict[str, Any], None]: + from openai import NOT_GIVEN + + usage_callback = kwargs.get("usage_callback") + completion_callback = kwargs.get("completion_callback") + + stream = await self.client.chat.completions.create( + model=model, + messages=messages, # type: ignore + stream=True, + stream_options={"include_usage": True}, + temperature=temperature if temperature is not None else NOT_GIVEN, + max_tokens=max_tokens if max_tokens is not None else NOT_GIVEN, + top_p=kwargs.get("top_p", NOT_GIVEN), + ) + + final_usage = None + + async for chunk in stream: + chunk_data = chunk.model_dump() + + if chunk.usage: + final_usage = chunk.usage.model_dump() + if usage_callback: + usage_callback(final_usage) + + yield chunk_data + + if completion_callback: + await completion_callback(model, final_usage) + + async def list_models(self) -> list[dict[str, Any]]: + """List available Gemini models.""" + try: + response = await self.client.models.list() + return [model.model_dump() for model in response.data] + except Exception as e: + from ...core.logging import get_logger + + logger = get_logger(__name__) + logger.error(f"Failed to list Gemini models: {e}") + return [] diff --git a/routstr/upstream/gemini.py b/routstr/upstream/gemini.py new file mode 100644 index 00000000..79a109c4 --- /dev/null +++ b/routstr/upstream/gemini.py @@ -0,0 +1,283 @@ +from __future__ import annotations + +import json +from collections.abc import AsyncGenerator +from typing import TYPE_CHECKING, Any + +from fastapi import Request +from fastapi.responses import Response, StreamingResponse + +from .base import BaseUpstreamProvider +from .clients.gemini import GeminiClient + +if TYPE_CHECKING: + from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow + from ..payment.models import Model + +from ..core.logging import get_logger + +logger = get_logger(__name__) + + +class GeminiUpstreamProvider(BaseUpstreamProvider): + provider_type = "gemini" + default_base_url = "https://generativelanguage.googleapis.com/v1beta" + platform_url = "https://aistudio.google.com/app/apikey" + + def __init__( + self, + base_url: str = "https://generativelanguage.googleapis.com/v1beta", + api_key: str = "", + provider_fee: float = 1.01, + ): + super().__init__( + api_key=api_key, + provider_fee=provider_fee, + base_url=base_url, + ) + self._client: GeminiClient | None = None + + @property + def client(self) -> GeminiClient: + """Get or create the Gemini API client.""" + if self._client is None: + self._client = GeminiClient(api_key=self.api_key) + return self._client + + @classmethod + def from_db_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "GeminiUpstreamProvider": + return cls( + base_url=provider_row.base_url, + api_key=provider_row.api_key, + provider_fee=provider_row.provider_fee, + ) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "Google Gemini", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + + def transform_model_name(self, model_id: str) -> str: + return model_id.removeprefix("gemini/") + + async def forward_request( + self, + request: Request, + path: str, + headers: dict, + request_body: bytes | None, + key: ApiKey, + max_cost_for_model: int, + session: AsyncSession, + model_obj: Model, + ) -> Response | StreamingResponse: + # Remove provider prefix from model ID for Gemini API + if "/" in model_obj.id: + model_obj.id = model_obj.id.split("/", 1)[1] + + if not path.startswith("chat/completions"): + return await super().forward_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + + if not request_body: + return await super().forward_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + + try: + openai_data = json.loads(request_body) + messages = openai_data.get("messages", []) + temperature = openai_data.get("temperature") + max_tokens = openai_data.get("max_tokens") + top_p = openai_data.get("top_p") + is_streaming = openai_data.get("stream", False) + + logger.info( + "Processing Gemini request with client abstraction", + extra={ + "model": model_obj.id, + "is_streaming": is_streaming, + "message_count": len(messages), + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + if is_streaming: + final_usage_data: dict | None = None + + def usage_callback(usage_data: dict[str, Any]) -> None: + """Callback to capture usage data during streaming""" + nonlocal final_usage_data + final_usage_data = usage_data + + async def completion_callback( + model: str, usage_data: dict[str, Any] | None + ) -> None: + """Callback to handle payment when streaming completes""" + nonlocal final_usage_data + if usage_data: + final_usage_data = usage_data + + payment_data = { + "model": model, + "usage": final_usage_data, + } + + from ..auth import adjust_payment_for_tokens + from ..core.db import create_session + + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if fresh_key: + try: + cost_data = await adjust_payment_for_tokens( + fresh_key, + payment_data, + new_session, + max_cost_for_model, + ) + + logger.info( + "Gemini streaming payment finalized", + extra={ + "cost_data": cost_data, + "usage_data": final_usage_data, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + except Exception as cost_error: + logger.error( + "Error finalizing Gemini streaming payment", + extra={ + "error": str(cost_error), + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + response_generator = self.client.generate_content_stream( + model=model_obj.id, + messages=messages, + temperature=temperature, + max_tokens=max_tokens, + top_p=top_p, + usage_callback=usage_callback, + completion_callback=completion_callback, + ) + + async def stream_with_cost() -> AsyncGenerator[bytes, None]: + try: + async for chunk in response_generator: + sse_data = f"data: {json.dumps(chunk)}\n\n" + yield sse_data.encode() + + except Exception as e: + logger.error( + "Error in Gemini streaming response", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + raise + + return StreamingResponse( + stream_with_cost(), + media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "Connection": "keep-alive"}, + ) + + else: + openai_format_response = await self.client.generate_content( + model=model_obj.id, + messages=messages, + temperature=temperature, + max_tokens=max_tokens, + top_p=top_p, + ) + + from ..auth import adjust_payment_for_tokens + + cost_data = await adjust_payment_for_tokens( + key, openai_format_response, session, max_cost_for_model + ) + openai_format_response["cost"] = cost_data + + logger.info( + "Gemini non-streaming payment completed", + extra={ + "cost_data": cost_data, + "model": model_obj.id, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + return Response( + content=json.dumps(openai_format_response), + media_type="application/json", + headers={"Cache-Control": "no-cache"}, + ) + + except Exception as e: + logger.error( + "Error in Gemini forward_request", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "path": path, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + return await super().forward_request( + request, + path, + headers, + request_body, + key, + max_cost_for_model, + session, + model_obj, + ) + + async def _fetch_provider_models(self) -> dict: + """Fetch models from Gemini API.""" + try: + models_data = await self.client.list_models() + + for model in models_data: + if "id" in model and model["id"].startswith("models/"): + model["id"] = model["id"].removeprefix("models/") + + return {"data": models_data} + except Exception as e: + logger.error( + f"Failed to fetch models from Gemini API: {e}", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "base_url": self.base_url, + }, + ) + return {"data": []} diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index 3d91550f..b453af1f 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -193,7 +193,7 @@ async def init_upstreams() -> list[BaseUpstreamProvider]: if provider: await provider.refresh_models_cache() upstreams.append(provider) - logger.info( + logger.debug( f"Initialized {provider_row.provider_type} provider", extra={ "base_url": provider_row.base_url, diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index cc0e2908..3ee58125 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -1,5 +1,7 @@ from typing import TYPE_CHECKING +import httpx + from ..payment.models import Model, async_fetch_openrouter_models from .base import BaseUpstreamProvider @@ -42,9 +44,33 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): "default_base_url": cls.default_base_url, "fixed_base_url": True, "platform_url": cls.platform_url, + "can_show_balance": True, } async def fetch_models(self) -> list[Model]: """Fetch all OpenRouter models.""" models_data = await async_fetch_openrouter_models() return [Model(**model) for model in models_data] # type: ignore + + async def get_balance(self) -> float | None: + """Get the current account balance from OpenRouter. + + Returns: + Float representing the balance amount (in credits/USD), or None if unavailable. + """ + url = f"{self.base_url}/credits" + headers = {"Authorization": f"Bearer {self.api_key}"} + + try: + async with httpx.AsyncClient(timeout=30.0) as client: + response = await client.get(url, headers=headers) + response.raise_for_status() + data = response.json() + + credits_data = data.get("data", {}) + total_credits = float(credits_data.get("total_credits", 0.0)) + total_usage = float(credits_data.get("total_usage", 0.0)) + + return total_credits - total_usage + except Exception: + return None diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index f8a9c965..0b4dfb08 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -36,6 +36,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): provider_type = "ppqai" default_base_url = "https://api.ppq.ai" platform_url = "https://ppq.ai/api-docs" + IGNORED_MODEL_IDS: list[str] = ["auto"] def __init__(self, api_key: str, provider_fee: float = 1.0): super().__init__( @@ -101,11 +102,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): url = f"{self.base_url}/models" headers = {"Authorization": f"Bearer {self.api_key}"} - logger.debug( - "Fetching models from PPQ.AI", - extra={"url": url, "has_api_key": bool(self.api_key)}, - ) - try: async with httpx.AsyncClient(timeout=30.0) as client: response = await client.get(url, headers=headers) @@ -113,10 +109,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): data = response.json() models_data = data.get("data", []) - logger.info( - "Fetched models from PPQ.AI", - extra={"model_count": len(models_data)}, - ) or_models = [ Model(**model) # type: ignore @@ -127,11 +119,16 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): for model_data in models_data: try: ppqai_model = PPQAIModel.parse_obj(model_data) + if ppqai_model.id in self.IGNORED_MODEL_IDS: + continue + or_model = next( ( model for model in or_models - if model.id == ppqai_model.id + if (model.id == ppqai_model.id) + or (model.id.split("/")[-1] == ppqai_model.id) + or (model.id == ppqai_model.id.split("/")[-1]) ), None, ) @@ -192,15 +189,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): return models - except httpx.HTTPStatusError as e: - logger.error( - "HTTP error fetching models from PPQ.AI", - extra={ - "status_code": e.response.status_code, - "error": str(e), - }, - ) - return [] except Exception as e: logger.error( "Error fetching models from PPQ.AI", @@ -371,16 +359,20 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): return topup_data - async def get_balance(self) -> dict[str, object]: + async def get_balance(self) -> float | None: """Get the current account balance from PPQ.AI. Returns: - Dict with balance information + Float representing the balance amount (in USD), or None if unavailable. Raises: httpx.HTTPStatusError: If the API request fails """ - return await self.check_balance() + data = await self.check_balance() + balance = data.get("balance") + if isinstance(balance, (int, float)): + return float(balance) + return None async def check_balance(self) -> dict[str, object]: """Check the account balance for this PPQ.AI account. diff --git a/tests/integration/test_performance_load.py b/tests/integration/test_performance_load.py index 581553b9..e5e5c3ce 100644 --- a/tests/integration/test_performance_load.py +++ b/tests/integration/test_performance_load.py @@ -142,46 +142,6 @@ class TestPerformanceBaseline: f" P99: {sorted(response_times)[int(len(response_times) * 0.99)]:.2f}ms" ) - @pytest.mark.asyncio - async def test_database_query_performance( - self, integration_session: Any, db_snapshot: Any - ) -> None: - """Test database operation performance""" - from sqlmodel import select - - from routstr.core.db import ApiKey - - # Create test data - for i in range(100): - key = ApiKey( - hashed_key=f"test_key_{i}", - balance=1000000, - total_spent=0, - total_requests=0, - ) - integration_session.add(key) - await integration_session.commit() - - # Test query performance - query_times = [] - - for _ in range(100): - start = time.time() - result = await integration_session.execute( - select(ApiKey).where(ApiKey.balance > 0) # type: ignore[arg-type] - ) - _ = result.all() - duration = (time.time() - start) * 1000 - query_times.append(duration) - - # All queries should complete < 100ms - assert max(query_times) < 100, ( - f"Max query time {max(query_times)}ms exceeds 100ms limit" - ) - print("\nDatabase query performance:") - print(f" Mean: {statistics.mean(query_times):.2f}ms") - print(f" Max: {max(query_times):.2f}ms") - @pytest.mark.integration @pytest.mark.slow diff --git a/ui/app/balances/page.tsx b/ui/app/balances/page.tsx index b5213287..89e44f3a 100644 --- a/ui/app/balances/page.tsx +++ b/ui/app/balances/page.tsx @@ -29,9 +29,7 @@ export default function BalancesPage() {
-

- Balances -

+

Balances

Monitor and manage wallet balances

diff --git a/ui/app/login/page.tsx b/ui/app/login/page.tsx index 7df50995..ba47a120 100644 --- a/ui/app/login/page.tsx +++ b/ui/app/login/page.tsx @@ -78,8 +78,8 @@ export default function AdminLoginPage(): ReactElement { }; return ( -
- +
+ Admin Login diff --git a/ui/app/logs/log-details-dialog.tsx b/ui/app/logs/log-details-dialog.tsx index 852aa4a8..5c0193ce 100644 --- a/ui/app/logs/log-details-dialog.tsx +++ b/ui/app/logs/log-details-dialog.tsx @@ -97,7 +97,7 @@ export function LogDetailsDialog({

Message

-
+                
                   {log.message}
                 
@@ -116,7 +116,12 @@ export function LogDetailsDialog({
- {totalModels === 0 ? ( - + {totalModels === 0 && ( +

No models found for this provider

-
-

Common issues:

-
    +
    +

    + Common issues: +

    +
    • - API credentials: Check if the API key is correct and has the right permissions + API credentials:{' '} + Check if the API key is correct + and has the right permissions
    • - Base URL: Verify the base URL is correct for your provider + Base URL: Verify + the base URL is correct for your + provider
    • - Network access: Ensure the server can reach the provider's API endpoint + Network access:{' '} + Ensure the server can reach the + provider's API endpoint
    • - Provider status: The upstream provider might be temporarily unavailable + Provider status:{' '} + The upstream provider might be + temporarily unavailable
    {groupData?.group_url && ( -

    - Current endpoint: {groupData.group_url} +

    + Current endpoint:{' '} + + {groupData.group_url} +

    )}
- ) : ( - )} +
); diff --git a/ui/app/page.tsx b/ui/app/page.tsx index 041512eb..ba5dff37 100644 --- a/ui/app/page.tsx +++ b/ui/app/page.tsx @@ -53,7 +53,7 @@ export default function DashboardPage() { window.removeEventListener('storage', syncAuthState); }; }, []); - + const { data: btcUsdPrice } = useQuery({ queryKey: ['btc-usd-price'], queryFn: fetchBtcUsdPrice, @@ -131,8 +131,13 @@ export default function DashboardPage() {
-

Dashboard

- +

+ Dashboard +

+
@@ -140,8 +145,9 @@ export default function DashboardPage() {

Usage Analytics

-

- Monitor requests, errors, and revenue over the last {timeRange} hours +

+ Monitor requests, errors, and revenue over the last {timeRange}{' '} + hours

@@ -176,27 +182,31 @@ export default function DashboardPage() {
{summaryLoading ? ( -
Loading summary...
+
Loading summary...
) : summaryData ? ( ) : null}
{metricsLoading ? ( -
+
Loading metrics...
) : metricsData && metricsData.metrics.length > 0 ? ( <> -
+
({ - ...m, - revenue_sats: m.revenue_msats / 1000, - refunds_sats: m.refunds_msats / 1000, - net_revenue_sats: - (m.revenue_msats - m.refunds_msats) / 1000, - })) as Array & { timestamp: string }>} + data={ + metricsData.metrics.map((m) => ({ + ...m, + revenue_sats: m.revenue_msats / 1000, + refunds_sats: m.refunds_msats / 1000, + net_revenue_sats: + (m.revenue_msats - m.refunds_msats) / 1000, + })) as Array< + Record & { timestamp: string } + > + } title='Revenue Over Time (sats)' dataKeys={[ { @@ -218,7 +228,11 @@ export default function DashboardPage() { />
& { timestamp: string }>} + data={ + metricsData.metrics as Array< + Record & { timestamp: string } + > + } title='Request Volume' dataKeys={[ { @@ -239,7 +253,11 @@ export default function DashboardPage() { ]} /> & { timestamp: string }>} + data={ + metricsData.metrics as Array< + Record & { timestamp: string } + > + } title='Error Tracking' dataKeys={[ { @@ -260,7 +278,11 @@ export default function DashboardPage() { ]} /> & { timestamp: string }>} + data={ + metricsData.metrics as Array< + Record & { timestamp: string } + > + } title='Payment Activity' dataKeys={[ { @@ -270,7 +292,7 @@ export default function DashboardPage() { }, ]} /> -
+
{summaryData && summaryData.unique_models.length > 0 && ( @@ -306,7 +328,9 @@ export default function DashboardPage() { key={type} className='flex items-center justify-between' > - {type} + + {type} + {count} @@ -335,16 +359,18 @@ export default function DashboardPage() {
{revenueByModelLoading ? ( -
Loading revenue by model...
+
+ Loading revenue by model... +
) : revenueByModelData && revenueByModelData.models.length > 0 ? ( - ) : null} {errorLoading ? ( -
Loading errors...
+
Loading errors...
) : errorData ? ( ) : null} diff --git a/ui/app/providers/page.tsx b/ui/app/providers/page.tsx index 186e5239..4624f26a 100644 --- a/ui/app/providers/page.tsx +++ b/ui/app/providers/page.tsx @@ -18,7 +18,9 @@ import { UpstreamProvider, CreateUpstreamProvider, UpdateUpstreamProvider, + AdminModel, } from '@/lib/api/services/admin'; +import { AddProviderModelDialog } from '@/components/AddProviderModelDialog'; import { Skeleton } from '@/components/ui/skeleton'; import { AlertCircle, @@ -54,7 +56,13 @@ import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs'; import { useState, useEffect } from 'react'; import { toast } from 'sonner'; -function ProviderBalance({ providerId }: { providerId: number }) { +function ProviderBalance({ + providerId, + platformUrl, +}: { + providerId: number; + platformUrl?: string | null; +}) { const [isTopupDialogOpen, setIsTopupDialogOpen] = useState(false); const [topupAmount, setTopupAmount] = useState(''); const [topupError, setTopupError] = useState(''); @@ -68,7 +76,11 @@ function ProviderBalance({ providerId }: { providerId: number }) { ); const queryClient = useQueryClient(); - const { data: balanceData, isLoading, error } = useQuery({ + const { + data: balanceData, + isLoading, + error, + } = useQuery({ queryKey: ['provider-balance', providerId], queryFn: () => AdminService.getProviderBalance(providerId), refetchInterval: 30000, @@ -100,7 +112,10 @@ function ProviderBalance({ providerId }: { providerId: number }) { mutationFn: async (amount: number) => { console.log('Calling top-up API with:', { providerId, amount }); try { - const result = await AdminService.initiateProviderTopup(providerId, amount); + const result = await AdminService.initiateProviderTopup( + providerId, + amount + ); console.log('API returned:', result); return result; } catch (err) { @@ -112,11 +127,8 @@ function ProviderBalance({ providerId }: { providerId: number }) { console.log('Top-up response:', data); console.log('Type of data:', typeof data); console.log('Keys in data:', Object.keys(data || {})); - - if ( - data?.topup_data?.payment_request && - data?.topup_data?.invoice_id - ) { + + if (data?.topup_data?.payment_request && data?.topup_data?.invoice_id) { setInvoiceData({ payment_request: data.topup_data.payment_request as string, invoice_id: data.topup_data.invoice_id as string, @@ -136,6 +148,10 @@ function ProviderBalance({ providerId }: { providerId: number }) { }); const handleTopup = () => { + // If no dialog open logic (which depends on API implementation), + // we check if we should redirect or open dialog based on available info + // But since this function is called inside the dialog, we might want to change + // how the "Top Up" button behaves instead. const amount = parseFloat(topupAmount); if (isNaN(amount)) { @@ -151,6 +167,35 @@ function ProviderBalance({ providerId }: { providerId: number }) { topupMutation.mutate(amount); }; + const handleTopUpClick = () => { + // Check if the provider supports direct topup (currently only PPQ.AI effectively) + // We can infer this if it's NOT OpenRouter or OpenAI, or strictly checking provider capability + // For now, we'll try to initiate topup for anyone, but if we know it fails (or isn't implemented), + // we should redirect. + // However, the prompt asks to redirect if topup is not implemented. + // The backend throws 500/400 if not implemented. + // A better approach is to check if we have a platform URL and maybe redirect there + // if we know it's not supported. + + // BUT, we don't know for sure if it's supported without checking metadata or trying. + // Let's rely on the "can_topup" metadata if available, but currently we only have "can_show_balance". + + // Simple heuristic: If platformUrl exists and we suspect no direct topup, redirect? + // Actually, let's try to open the dialog, but if it's OpenRouter/OpenAI, maybe we just redirect? + // The user specifically mentioned "like in openrouter". + + if ( + platformUrl && + (platformUrl.includes('openrouter.ai') || + platformUrl.includes('openai.com')) + ) { + window.open(platformUrl, '_blank'); + return; + } + + setIsTopupDialogOpen(true); + }; + const handleCloseDialog = () => { setIsTopupDialogOpen(false); setTopupAmount(''); @@ -170,12 +215,18 @@ function ProviderBalance({ providerId }: { providerId: number }) { const balance = balanceData.balance_data; let displayValue = 'N/A'; - if (typeof balance.balance === 'number') { - displayValue = `$${balance.balance.toFixed(2)}`; - } else if (typeof balance.balance === 'string') { - displayValue = balance.balance; - } else if (balance.amount !== undefined) { - displayValue = `$${Number(balance.amount).toFixed(2)}`; + if (typeof balance === 'number') { + displayValue = `$${balance.toFixed(2)}`; + } else if (balance && typeof balance === 'object') { + // Legacy support for object response + const b = balance as Record; + if (typeof b.balance === 'number') { + displayValue = `$${b.balance.toFixed(2)}`; + } else if (typeof b.balance === 'string') { + displayValue = b.balance; + } else if (b.amount !== undefined) { + displayValue = `$${Number(b.amount).toFixed(2)}`; + } } return ( @@ -183,7 +234,7 @@ function ProviderBalance({ providerId }: { providerId: number }) {
-

- Top-up successful! -

+

Top-up successful!

) : invoiceData ? (
@@ -338,6 +387,17 @@ export default function ProvidersPage() { ); const [viewingModels, setViewingModels] = useState(null); const [isCreatingAccount, setIsCreatingAccount] = useState(false); + const [modelDialogState, setModelDialogState] = useState<{ + isOpen: boolean; + providerId: number | null; + mode: 'create' | 'edit' | 'override'; + initialData?: AdminModel | null; + }>({ + isOpen: false, + providerId: null, + mode: 'create', + initialData: null, + }); const [formData, setFormData] = useState({ provider_type: 'openrouter', @@ -428,7 +488,9 @@ export default function ProvidersPage() { ...formData, api_key: String(response.account_data.api_key), }); - toast.success('Account created successfully! API key has been filled in.'); + toast.success( + 'Account created successfully! API key has been filled in.' + ); } else { toast.success('Account created, but no API key returned.'); } @@ -533,6 +595,33 @@ export default function ProvidersPage() { } }; + const handleAddModel = (providerId: number) => { + setModelDialogState({ + isOpen: true, + providerId, + mode: 'create', + initialData: null, + }); + }; + + const handleEditModel = (providerId: number, model: AdminModel) => { + setModelDialogState({ + isOpen: true, + providerId, + mode: 'edit', + initialData: model, + }); + }; + + const handleOverrideModel = (providerId: number, model: AdminModel) => { + setModelDialogState({ + isOpen: true, + providerId, + mode: 'override', + initialData: model, + }); + }; + return ( @@ -697,8 +786,7 @@ export default function ProvidersPage() { )} />

- 1.01 means +1% e.g. currency exchange, card - fees, etc. + 1.01 means +1% e.g. currency exchange, card fees, etc.

@@ -775,7 +863,12 @@ export default function ProvidersPage() {
{canShowBalance(provider.provider_type) && provider.api_key && ( - + )}
) : (
@@ -872,9 +977,24 @@ export default function ProvidersPage() { {model.description || model.name}
-
- {model.context_length?.toLocaleString()}{' '} - tokens +
+
+ {model.context_length?.toLocaleString()}{' '} + tokens +
+
))} @@ -925,12 +1045,24 @@ export default function ProvidersPage() { value='custom' className='mt-4 space-y-2' > - {providerModels.db_models.length > 0 && ( -
- Custom models override or extend the - provider's catalog. -
- )} +
+ {providerModels.db_models.length > 0 && ( +
+ Custom models override or extend the + provider's catalog. +
+ )} + +
{providerModels.db_models.length === 0 ? (
No custom models configured @@ -966,9 +1098,24 @@ export default function ProvidersPage() { model.name}
-
- {model.context_length?.toLocaleString()}{' '} - tokens +
+
+ {model.context_length?.toLocaleString()}{' '} + tokens +
+
) @@ -1003,9 +1150,25 @@ export default function ProvidersPage() { model.name}
-
- {model.context_length?.toLocaleString()}{' '} - tokens +
+
+ {model.context_length?.toLocaleString()}{' '} + tokens +
+
) @@ -1120,7 +1283,7 @@ export default function ProvidersPage() {
)}
- @@ -1152,8 +1315,7 @@ export default function ProvidersPage() { )} />

- 1.01 means +1% e.g. currency exchange, card - fees, etc. + 1.01 means +1% e.g. currency exchange, card fees, etc.

@@ -1173,6 +1335,23 @@ export default function ProvidersPage() { + + {modelDialogState.providerId && ( + + setModelDialogState((prev) => ({ ...prev, isOpen: false })) + } + onSuccess={() => { + queryClient.invalidateQueries({ + queryKey: ['provider-models', modelDialogState.providerId], + }); + }} + initialData={modelDialogState.initialData} + mode={modelDialogState.mode} + /> + )} ); diff --git a/ui/components/AddProviderModelDialog.tsx b/ui/components/AddProviderModelDialog.tsx new file mode 100644 index 00000000..301682f4 --- /dev/null +++ b/ui/components/AddProviderModelDialog.tsx @@ -0,0 +1,956 @@ +'use client'; + +import React, { useEffect, useMemo, useState } from 'react'; +import { useForm } from 'react-hook-form'; +import { z } from 'zod'; +import { zodResolver } from '@hookform/resolvers/zod'; +import { useQuery } from '@tanstack/react-query'; +import { Button } from '@/components/ui/button'; +import { Input } from '@/components/ui/input'; +import { Textarea } from '@/components/ui/textarea'; +import { + Popover, + PopoverContent, + PopoverTrigger, +} from '@/components/ui/popover'; +import { + Command, + CommandEmpty, + CommandGroup, + CommandInput, + CommandItem, + CommandList, +} from '@/components/ui/command'; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from '@/components/ui/dialog'; +import { + Form, + FormControl, + FormDescription, + FormField, + FormItem, + FormLabel, + FormMessage, +} from '@/components/ui/form'; +import { Switch } from '@/components/ui/switch'; +import { Loader2, Plus } from 'lucide-react'; +import { toast } from 'sonner'; +import { AdminService, type AdminModel } from '@/lib/api/services/admin'; + +const listFromString = (value: string): string[] => + value + .split(',') + .map((item) => item.trim()) + .filter((item) => item.length > 0); + +const listToString = (value: string[] | undefined | null): string => + value && value.length > 0 ? value.join(', ') : ''; + +const FormSchema = z.object({ + id: z.string().min(1, 'Model ID is required'), + name: z.string().min(1, 'Name is required'), + description: z.string().default(''), + context_length: z.coerce.number().min(0).default(8192), + modality: z.string().min(1, 'Modality is required'), + input_modalities_raw: z.string().default(''), + output_modalities_raw: z.string().default(''), + tokenizer: z.string().default(''), + instruct_type: z.string().default(''), + canonical_slug: z.string().default(''), + alias_ids_raw: z.string().default(''), + upstream_provider_id: z.string().default(''), + input_cost: z.coerce.number().min(0).default(0), + output_cost: z.coerce.number().min(0).default(0), + request_cost: z.coerce.number().min(0).default(0), + image_cost: z.coerce.number().min(0).default(0), + web_search_cost: z.coerce.number().min(0).default(0), + internal_reasoning_cost: z.coerce.number().min(0).default(0), + max_prompt_cost: z.coerce.number().min(0).default(0), + max_completion_cost: z.coerce.number().min(0).default(0), + max_cost: z.coerce.number().min(0).default(0), + per_request_limits_raw: z.string().default(''), + top_provider_context_length: z.coerce.number().min(0).optional(), + top_provider_max_completion_tokens: z.coerce.number().min(0).optional(), + top_provider_is_moderated: z.boolean().default(false), + enabled: z.boolean().default(true), +}); + +type FormData = z.output; + +export interface AddProviderModelDialogProps { + providerId: number; + isOpen: boolean; + onClose: () => void; + onSuccess: () => void; + initialData?: AdminModel | null; + mode?: 'create' | 'edit' | 'override'; +} + +export function AddProviderModelDialog({ + providerId, + isOpen, + onClose, + onSuccess, + initialData, + mode = 'create', +}: AddProviderModelDialogProps) { + const [isSubmitting, setIsSubmitting] = useState(false); + const [isPresetOpen, setIsPresetOpen] = useState(false); + const [selectedPresetLabel, setSelectedPresetLabel] = + useState('Select a preset'); + + const form = useForm({ + resolver: zodResolver(FormSchema) as never, + defaultValues: { + id: '', + name: '', + description: '', + context_length: 8192, + modality: 'text', + input_modalities_raw: 'text', + output_modalities_raw: 'text', + tokenizer: '', + instruct_type: '', + canonical_slug: '', + alias_ids_raw: '', + upstream_provider_id: '', + input_cost: 0, + output_cost: 0, + request_cost: 0, + image_cost: 0, + web_search_cost: 0, + internal_reasoning_cost: 0, + max_prompt_cost: 0, + max_completion_cost: 0, + max_cost: 0, + per_request_limits_raw: '', + top_provider_context_length: undefined, + top_provider_max_completion_tokens: undefined, + top_provider_is_moderated: false, + enabled: true, + }, + }); + + const isOverride = useMemo(() => mode === 'override', [mode]); + const isEdit = useMemo(() => mode === 'edit', [mode]); + const { data: presets = [], isLoading: isLoadingPresets } = useQuery({ + queryKey: ['openrouter-presets'], + queryFn: () => AdminService.getOpenRouterPresets(), + staleTime: 10 * 60 * 1000, + refetchOnWindowFocus: false, + }); + + useEffect(() => { + if (initialData) { + const architecture = initialData.architecture as Record; + const pricing = initialData.pricing as Record; + const topProvider = initialData.top_provider as Record< + string, + unknown + > | null; + + form.reset({ + id: initialData.id, + name: initialData.name, + description: initialData.description, + context_length: initialData.context_length, + modality: + typeof architecture?.modality === 'string' + ? architecture.modality + : 'text', + input_modalities_raw: listToString( + (architecture?.input_modalities as string[]) || [] + ), + output_modalities_raw: listToString( + (architecture?.output_modalities as string[]) || [] + ), + tokenizer: + typeof architecture?.tokenizer === 'string' + ? architecture.tokenizer + : '', + instruct_type: + typeof architecture?.instruct_type === 'string' + ? architecture.instruct_type + : '', + canonical_slug: initialData.canonical_slug || '', + alias_ids_raw: listToString(initialData.alias_ids), + upstream_provider_id: + typeof initialData.upstream_provider_id === 'string' + ? initialData.upstream_provider_id + : initialData.upstream_provider_id?.toString() || '', + input_cost: pricing?.prompt ?? 0, + output_cost: pricing?.completion ?? 0, + request_cost: pricing?.request ?? 0, + image_cost: pricing?.image ?? 0, + web_search_cost: pricing?.web_search ?? 0, + internal_reasoning_cost: pricing?.internal_reasoning ?? 0, + max_prompt_cost: pricing?.max_prompt_cost ?? 0, + max_completion_cost: pricing?.max_completion_cost ?? 0, + max_cost: pricing?.max_cost ?? 0, + per_request_limits_raw: initialData.per_request_limits + ? JSON.stringify(initialData.per_request_limits, null, 2) + : '', + top_provider_context_length: + typeof topProvider?.context_length === 'number' + ? topProvider.context_length + : undefined, + top_provider_max_completion_tokens: + typeof topProvider?.max_completion_tokens === 'number' + ? topProvider.max_completion_tokens + : undefined, + top_provider_is_moderated: + typeof topProvider?.is_moderated === 'boolean' + ? topProvider.is_moderated + : false, + enabled: initialData.enabled, + }); + } else { + form.reset({ + id: '', + name: '', + description: '', + context_length: 8192, + modality: 'text', + input_modalities_raw: 'text', + output_modalities_raw: 'text', + tokenizer: '', + instruct_type: '', + canonical_slug: '', + alias_ids_raw: '', + upstream_provider_id: '', + input_cost: 0, + output_cost: 0, + request_cost: 0, + image_cost: 0, + web_search_cost: 0, + internal_reasoning_cost: 0, + max_prompt_cost: 0, + max_completion_cost: 0, + max_cost: 0, + per_request_limits_raw: '', + top_provider_context_length: undefined, + top_provider_max_completion_tokens: undefined, + top_provider_is_moderated: false, + enabled: true, + }); + } + }, [initialData, form, isOpen]); + + const applyModelToForm = (model: AdminModel) => { + setSelectedPresetLabel(`${model.id} — ${model.name}`); + const architecture = model.architecture as Record; + const pricing = model.pricing as Record; + const topProvider = model.top_provider as Record | null; + + form.setValue('id', model.id); + form.setValue('name', model.name); + form.setValue('description', model.description || ''); + form.setValue('context_length', model.context_length); + form.setValue( + 'modality', + typeof architecture?.modality === 'string' + ? architecture.modality + : 'text' + ); + form.setValue( + 'input_modalities_raw', + listToString((architecture?.input_modalities as string[]) || []) + ); + form.setValue( + 'output_modalities_raw', + listToString((architecture?.output_modalities as string[]) || []) + ); + form.setValue( + 'tokenizer', + typeof architecture?.tokenizer === 'string' ? architecture.tokenizer : '' + ); + form.setValue( + 'instruct_type', + typeof architecture?.instruct_type === 'string' + ? architecture.instruct_type + : '' + ); + form.setValue('canonical_slug', model.canonical_slug || ''); + form.setValue('alias_ids_raw', listToString(model.alias_ids)); + form.setValue( + 'upstream_provider_id', + typeof model.upstream_provider_id === 'string' + ? model.upstream_provider_id + : model.upstream_provider_id?.toString() || '' + ); + form.setValue('input_cost', pricing?.prompt ?? 0); + form.setValue('output_cost', pricing?.completion ?? 0); + form.setValue('request_cost', pricing?.request ?? 0); + form.setValue('image_cost', pricing?.image ?? 0); + form.setValue('web_search_cost', pricing?.web_search ?? 0); + form.setValue('internal_reasoning_cost', pricing?.internal_reasoning ?? 0); + form.setValue('max_prompt_cost', pricing?.max_prompt_cost ?? 0); + form.setValue('max_completion_cost', pricing?.max_completion_cost ?? 0); + form.setValue('max_cost', pricing?.max_cost ?? 0); + form.setValue( + 'per_request_limits_raw', + model.per_request_limits + ? JSON.stringify(model.per_request_limits, null, 2) + : '' + ); + form.setValue( + 'top_provider_context_length', + typeof topProvider?.context_length === 'number' + ? topProvider.context_length + : undefined + ); + form.setValue( + 'top_provider_max_completion_tokens', + typeof topProvider?.max_completion_tokens === 'number' + ? topProvider.max_completion_tokens + : undefined + ); + form.setValue( + 'top_provider_is_moderated', + typeof topProvider?.is_moderated === 'boolean' + ? topProvider.is_moderated + : false + ); + form.setValue('enabled', model.enabled); + }; + + const onSubmit = async (data: FormData) => { + setIsSubmitting(true); + try { + let perRequestLimits: Record | null = null; + if ( + data.per_request_limits_raw && + data.per_request_limits_raw.trim().length + ) { + try { + perRequestLimits = JSON.parse(data.per_request_limits_raw); + } catch { + toast.error('Per-request limits must be valid JSON'); + setIsSubmitting(false); + return; + } + } + + const adminModel: AdminModel = { + id: data.id, + name: data.name, + description: data.description || '', + created: Math.floor(Date.now() / 1000), + context_length: data.context_length, + architecture: { + modality: data.modality, + input_modalities: listFromString( + data.input_modalities_raw || data.modality + ), + output_modalities: listFromString( + data.output_modalities_raw || data.modality + ), + tokenizer: data.tokenizer || '', + instruct_type: data.instruct_type?.trim() || null, + }, + pricing: { + prompt: data.input_cost, + completion: data.output_cost, + request: data.request_cost, + image: data.image_cost, + web_search: data.web_search_cost, + internal_reasoning: data.internal_reasoning_cost, + max_prompt_cost: data.max_prompt_cost, + max_completion_cost: data.max_completion_cost, + max_cost: data.max_cost, + }, + per_request_limits: perRequestLimits, + top_provider: + data.top_provider_context_length || + data.top_provider_max_completion_tokens || + data.top_provider_is_moderated + ? { + context_length: data.top_provider_context_length ?? null, + max_completion_tokens: + data.top_provider_max_completion_tokens ?? null, + is_moderated: data.top_provider_is_moderated, + } + : null, + upstream_provider_id: data.upstream_provider_id?.trim().length + ? data.upstream_provider_id.trim() + : providerId, + canonical_slug: data.canonical_slug?.trim() || null, + alias_ids: listFromString(data.alias_ids_raw || ''), + enabled: data.enabled, + }; + + if (isEdit) { + await AdminService.updateProviderModel(providerId, data.id, adminModel); + toast.success('Model updated successfully'); + } else { + await AdminService.createProviderModel(providerId, adminModel); + toast.success( + isOverride ? 'Model override created' : 'Model created successfully' + ); + } + + onSuccess(); + onClose(); + } catch (error: unknown) { + const message = + error instanceof Error ? error.message : 'Unknown error saving model'; + toast.error(`Failed to save model: ${message}`); + } finally { + setIsSubmitting(false); + } + }; + + const title = isEdit + ? 'Edit Model' + : isOverride + ? 'Override Model' + : 'Add Custom Model'; + const description = isOverride + ? 'Create a custom override for this upstream model' + : 'Add a new model configuration for this provider'; + + return ( + !open && onClose()}> + + + + + {title} + + {description} + + {!isEdit && !isOverride && ( +
+
Presets
+
+ + + + + e.preventDefault()} + > + + + e.stopPropagation()} + > + {isLoadingPresets ? ( + Loading presets... + ) : presets.length === 0 ? ( + No presets available. + ) : ( + + {presets.map((preset) => ( + { + applyModelToForm(preset); + setIsPresetOpen(false); + }} + > +
+ {preset.id} + + {preset.name} + +
+
+ ))} +
+ )} +
+
+
+
+
+
+ Prefill fields from a preset model definition, then adjust as + needed. +
+
+ )} +
+ +
+ ( + + Model ID * + + + + + Unique identifier for the model + + + + )} + /> + + ( + + Display Name * + + + + + + )} + /> +
+ + ( + + Description + +