diff --git a/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py b/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py new file mode 100644 index 00000000..c140b557 --- /dev/null +++ b/migrations/versions/c6d7e8f9a0b1_add_slug_to_upstream_providers.py @@ -0,0 +1,88 @@ +"""add slug to upstream_providers + +Revision ID: c6d7e8f9a0b1 +Revises: b5e7c9d1f3a2 +Create Date: 2026-06-29 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +from routstr.core.provider_slugs import provider_slug_base, provider_slug_candidate + +revision = "c6d7e8f9a0b1" +down_revision = "b5e7c9d1f3a2" +branch_labels = None +depends_on = None + + +def _allocate_backfill_slug(provider_type: str, reserved_slugs: set[str]) -> str: + base = provider_slug_base(provider_type) + suffix_number = 1 + while True: + candidate = provider_slug_candidate(base, suffix_number) + if candidate not in reserved_slugs: + reserved_slugs.add(candidate) + return candidate + suffix_number += 1 + + +def _backfill_provider_slugs(conn: sa.Connection) -> None: + existing_rows = conn.execute( + sa.text( + "SELECT slug FROM upstream_providers " + "WHERE slug IS NOT NULL AND slug != ''" + ) + ) + reserved_slugs = {str(row.slug).lower() for row in existing_rows} + + rows_to_backfill = conn.execute( + sa.text( + "SELECT id, provider_type FROM upstream_providers " + "WHERE slug IS NULL OR slug = '' " + "ORDER BY id" + ) + ) + for row in rows_to_backfill: + slug = _allocate_backfill_slug(str(row.provider_type), reserved_slugs) + conn.execute( + sa.text("UPDATE upstream_providers SET slug = :slug WHERE id = :id"), + {"slug": slug, "id": row.id}, + ) + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + columns = {c["name"] for c in inspector.get_columns("upstream_providers")} + + if "slug" not in columns: + op.add_column( + "upstream_providers", + sa.Column("slug", sa.String(), nullable=True), + ) + + _backfill_provider_slugs(conn) + + existing_indexes = {idx["name"] for idx in inspector.get_indexes("upstream_providers")} + if "ix_upstream_providers_slug" not in existing_indexes: + op.create_index( + "ix_upstream_providers_slug", + "upstream_providers", + ["slug"], + unique=True, + ) + + +def downgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + existing_indexes = {idx["name"] for idx in inspector.get_indexes("upstream_providers")} + if "ix_upstream_providers_slug" in existing_indexes: + op.drop_index("ix_upstream_providers_slug", table_name="upstream_providers") + + columns = {c["name"] for c in inspector.get_columns("upstream_providers")} + if "slug" in columns: + op.drop_column("upstream_providers", "slug") diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 3fffc8c7..b0fc080e 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1,5 +1,6 @@ import asyncio import json +import re import secrets from datetime import datetime, timezone from pathlib import Path @@ -8,6 +9,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request from pydantic import BaseModel, RootModel from pydantic.v1 import ValidationError as PydanticValidationError from sqlmodel import select +from sqlmodel.ext.asyncio.session import AsyncSession from ..payment.models import _row_to_model, list_models from ..proxy import refresh_model_maps, reinitialize_upstreams @@ -29,6 +31,7 @@ from .db import ( ) from .log_manager import log_manager from .logging import get_logger +from .provider_slugs import allocate_unique_provider_slug from .settings import SettingsService, settings logger = get_logger(__name__) @@ -456,19 +459,18 @@ class ModelCreate(BaseModel): dependencies=[Depends(require_admin_api)], ) async def upsert_provider_model( - provider_id: int, payload: ModelCreate + provider_id: str, payload: ModelCreate ) -> dict[str, object]: print(payload) logger.info( f"UPSERT_PROVIDER_MODEL called: provider_id={provider_id}, model_id={payload.id}" ) 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 = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) # Try to get existing model - existing_row = await session.get(ModelRow, (payload.id, provider_id)) + existing_row = await session.get(ModelRow, (payload.id, provider_pk)) if existing_row: # Update existing model @@ -524,7 +526,7 @@ async def upsert_provider_model( alias_ids=( json.dumps(payload.alias_ids) if payload.alias_ids else None ), - upstream_provider_id=provider_id, + upstream_provider_id=provider_pk, enabled=payload.enabled, forwarded_model_id=payload.forwarded_model_id or payload.id, ) @@ -543,7 +545,7 @@ async def upsert_provider_model( dependencies=[Depends(require_admin_api)], ) async def update_provider_model_legacy( - provider_id: int, model_id: str, payload: ModelCreate + provider_id: str, model_id: str, payload: ModelCreate ) -> dict[str, object]: """Legacy PATCH endpoint - redirects to upsert POST endpoint for backward compatibility.""" logger.info( @@ -556,13 +558,12 @@ async def update_provider_model_legacy( "/api/upstream-providers/{provider_id}/models/{model_id:path}", dependencies=[Depends(require_admin_api)], ) -async def get_provider_model(provider_id: int, model_id: str) -> dict[str, object]: +async def get_provider_model(provider_id: str, model_id: str) -> dict[str, object]: 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 = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) - row = await session.get(ModelRow, (model_id, provider_id)) + row = await session.get(ModelRow, (model_id, provider_pk)) if not row: raise HTTPException( status_code=404, detail="Model not found for this provider" @@ -576,9 +577,11 @@ async def get_provider_model(provider_id: int, model_id: str) -> dict[str, objec "/api/upstream-providers/{provider_id}/models/{model_id:path}", dependencies=[Depends(require_admin_api)], ) -async def delete_provider_model(provider_id: int, model_id: str) -> dict[str, object]: +async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, object]: async with create_session() as session: - row = await session.get(ModelRow, (model_id, provider_id)) + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) + row = await session.get(ModelRow, (model_id, provider_pk)) if not row: raise HTTPException( status_code=404, detail="Model not found for this provider" @@ -593,10 +596,12 @@ async def delete_provider_model(provider_id: int, model_id: str) -> dict[str, ob "/api/upstream-providers/{provider_id}/models", dependencies=[Depends(require_admin_api)], ) -async def delete_all_provider_models(provider_id: int) -> dict[str, object]: +async def delete_all_provider_models(provider_id: str) -> dict[str, object]: async with create_session() as session: + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) result = await session.exec( - select(ModelRow).where(ModelRow.upstream_provider_id == provider_id) + select(ModelRow).where(ModelRow.upstream_provider_id == provider_pk) ) # type: ignore rows = result.all() for row in rows: @@ -615,7 +620,7 @@ class BatchOverrideRequest(BaseModel): dependencies=[Depends(require_admin_api)], ) async def batch_override_provider_models( - provider_id: int, payload: BatchOverrideRequest + provider_id: str, payload: BatchOverrideRequest ) -> dict[str, object]: """Batch override models for a specific provider.""" logger.info( @@ -623,15 +628,14 @@ async def batch_override_provider_models( ) 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 = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) overridden_count = 0 for model_data in payload.models: # Try to get existing model regardless of whether it's enabled or not - existing_row = await session.get(ModelRow, (model_data.id, provider_id)) + existing_row = await session.get(ModelRow, (model_data.id, provider_pk)) if existing_row: # Update existing @@ -685,7 +689,7 @@ async def batch_override_provider_models( if model_data.alias_ids else None ), - upstream_provider_id=provider_id, + upstream_provider_id=provider_pk, enabled=model_data.enabled, ) session.add(row) @@ -702,6 +706,85 @@ async def batch_override_provider_models( } +_SLUG_PATTERN = re.compile(r"^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$") + + +def _validate_slug(value: str) -> str: + candidate = value.strip().lower() + if not _SLUG_PATTERN.fullmatch(candidate): + raise HTTPException( + status_code=400, + detail=( + "slug must be 3-64 chars, lowercase letters/digits/hyphens, " + "and may not start or end with a hyphen" + ), + ) + if candidate.isdigit(): + raise HTTPException( + status_code=400, + detail="slug must not be all digits", + ) + return candidate + + +async def _ensure_unique_slug( + session: AsyncSession, slug: str, exclude_id: int | None = None +) -> None: + stmt = select(UpstreamProviderRow).where(UpstreamProviderRow.slug == slug) + result = await session.exec(stmt) + existing = result.first() + if existing and existing.id != exclude_id: + raise HTTPException( + status_code=409, + detail="Provider with this slug already exists", + ) + + +async def _get_upstream_provider_by_ref( + session: AsyncSession, provider_ref: str +) -> UpstreamProviderRow: + if provider_ref.isdigit(): + provider = await session.get(UpstreamProviderRow, int(provider_ref)) + else: + slug = _validate_slug(provider_ref) + result = await session.exec( + select(UpstreamProviderRow).where(UpstreamProviderRow.slug == slug) + ) + provider = result.first() + + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + return provider + + +def _provider_pk(provider: UpstreamProviderRow) -> int: + if provider.id is None: + raise HTTPException(status_code=500, detail="Provider has no database id") + return provider.id + + +def _serialize_provider( + provider: UpstreamProviderRow, redact_api_key: bool = True +) -> dict[str, object]: + return { + "id": provider.id, + "slug": provider.slug, + "provider_type": provider.provider_type, + "base_url": provider.base_url, + "api_key": "[REDACTED]" + if (redact_api_key and provider.api_key) + else provider.api_key + if not redact_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, + } + + class UpstreamProviderCreate(BaseModel): provider_type: str base_url: str @@ -710,6 +793,7 @@ class UpstreamProviderCreate(BaseModel): enabled: bool = True provider_fee: float = 1.01 provider_settings: dict | None = None + slug: str | None = None class UpstreamProviderUpdate(BaseModel): @@ -720,6 +804,50 @@ class UpstreamProviderUpdate(BaseModel): enabled: bool | None = None provider_fee: float | None = None provider_settings: dict | None = None + slug: str | None = None + + +class UpstreamProviderUpdateBySlug(BaseModel): + slug: str + new_slug: str | None = None + provider_type: str | None = None + base_url: str | None = None + api_key: str | None = None + api_version: str | None = None + enabled: bool | None = None + provider_fee: float | None = None + provider_settings: dict | None = None + + +async def _apply_provider_update( + session: AsyncSession, + provider: UpstreamProviderRow, + payload: UpstreamProviderUpdate, + new_slug: str | None = None, +) -> None: + if new_slug is not None: + validated = _validate_slug(new_slug) + await _ensure_unique_slug(session, validated, exclude_id=provider.id) + provider.slug = validated + + if payload.provider_type is not None: + provider.provider_type = payload.provider_type + if payload.base_url is not None: + provider.base_url = payload.base_url + if payload.api_key is not None: + provider.api_key = payload.api_key + if payload.api_version is not None: + provider.api_version = payload.api_version + if payload.enabled is not None: + provider.enabled = payload.enabled + if payload.provider_fee is not None: + provider.provider_fee = payload.provider_fee + if payload.provider_settings is not None: + provider.provider_settings = json.dumps(payload.provider_settings) + + session.add(provider) + await session.commit() + await session.refresh(provider) @admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) @@ -727,21 +855,7 @@ 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 [_serialize_provider(p) for p in providers] @admin_router.post("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) @@ -761,7 +875,14 @@ async def create_upstream_provider( detail="Provider with this base URL and API key already exists", ) + if payload.slug: + slug = _validate_slug(payload.slug) + await _ensure_unique_slug(session, slug) + else: + slug = await allocate_unique_provider_slug(session, payload.provider_type) + provider = UpstreamProviderRow( + slug=slug, provider_type=payload.provider_type, base_url=payload.base_url, api_key=payload.api_key, @@ -778,99 +899,81 @@ async def create_upstream_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 _serialize_provider(provider) @admin_router.get( "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] ) -async def get_upstream_provider(provider_id: int) -> dict[str, object]: +async def get_upstream_provider(provider_id: str) -> dict[str, object]: 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 { - "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, - } + provider = await _get_upstream_provider_by_ref(session, provider_id) + return _serialize_provider(provider) @admin_router.patch( "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] ) async def update_upstream_provider( - provider_id: int, payload: UpstreamProviderUpdate + provider_id: str, payload: UpstreamProviderUpdate ) -> dict[str, object]: 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 = await _get_upstream_provider_by_ref(session, provider_id) - if payload.provider_type is not None: - provider.provider_type = payload.provider_type - if payload.base_url is not None: - provider.base_url = payload.base_url - if payload.api_key is not None: - provider.api_key = payload.api_key - if payload.api_version is not None: - provider.api_version = payload.api_version - if payload.enabled is not None: - provider.enabled = payload.enabled - if payload.provider_fee is not None: - provider.provider_fee = payload.provider_fee - if payload.provider_settings is not None: - provider.provider_settings = json.dumps(payload.provider_settings) - - session.add(provider) - await session.commit() - await session.refresh(provider) + await _apply_provider_update(session, provider, payload, new_slug=payload.slug) 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 _serialize_provider(provider) + + +@admin_router.patch( + "/api/upstream-providers", dependencies=[Depends(require_admin_api)] +) +async def update_upstream_provider_by_slug( + payload: UpstreamProviderUpdateBySlug, +) -> dict[str, object]: + lookup = _validate_slug(payload.slug) + async with create_session() as session: + result = await session.exec( + select(UpstreamProviderRow).where( + UpstreamProviderRow.slug == lookup + ) + ) + provider = result.first() + if not provider: + raise HTTPException(status_code=404, detail="Provider not found") + + update_payload = UpstreamProviderUpdate( + provider_type=payload.provider_type, + base_url=payload.base_url, + api_key=payload.api_key, + api_version=payload.api_version, + enabled=payload.enabled, + provider_fee=payload.provider_fee, + provider_settings=payload.provider_settings, + ) + await _apply_provider_update( + session, provider, update_payload, new_slug=payload.new_slug + ) + + await reinitialize_upstreams() + await refresh_model_maps() + return _serialize_provider(provider) @admin_router.delete( "/api/upstream-providers/{provider_id}", dependencies=[Depends(require_admin_api)] ) -async def delete_upstream_provider(provider_id: int) -> dict[str, object]: +async def delete_upstream_provider(provider_id: str) -> dict[str, object]: 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 = await _get_upstream_provider_by_ref(session, provider_id) + deleted_id = _provider_pk(provider) await session.delete(provider) await session.commit() await reinitialize_upstreams() await refresh_model_maps() - return {"ok": True, "deleted_id": provider_id} + return {"ok": True, "deleted_id": deleted_id} @admin_router.get("/api/provider-types", dependencies=[Depends(require_admin_api)]) @@ -885,17 +988,16 @@ async def get_provider_types() -> list[dict[str, object]]: "/api/upstream-providers/{provider_id}/models", dependencies=[Depends(require_admin_api)], ) -async def get_provider_models(provider_id: int) -> dict[str, object]: +async def get_provider_models(provider_id: str) -> dict[str, object]: from ..upstream.helpers import _instantiate_provider 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 = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) db_models = await list_models( session=session, - upstream_id=provider_id, + upstream_id=provider_pk, include_disabled=True, apply_fees=False, ) @@ -985,13 +1087,11 @@ class TopupTokenRequest(BaseModel): dependencies=[Depends(require_admin_api)], ) async def topup_provider_with_token( - provider_id: int, payload: TopupTokenRequest + provider_id: str, payload: TopupTokenRequest ) -> dict: """Redeem a Cashu token for an upstream provider.""" 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 = await _get_upstream_provider_by_ref(session, provider_id) import httpx @@ -1022,15 +1122,13 @@ async def topup_provider_with_token( dependencies=[Depends(require_admin_api)], ) async def initiate_provider_topup( - provider_id: int, payload: TopupRequest + provider_id: str, payload: TopupRequest ) -> dict[str, object]: """Initiate a Lightning Network top-up for the upstream provider account.""" from ..upstream.helpers import _instantiate_provider 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 = await _get_upstream_provider_by_ref(session, provider_id) try: logger.info( @@ -1150,15 +1248,13 @@ async def initiate_provider_topup( "/api/upstream-providers/{provider_id}/topup/{invoice_id}/status", dependencies=[Depends(require_admin_api)], ) -async def check_topup_status(provider_id: int, invoice_id: str) -> dict[str, object]: +async def check_topup_status(provider_id: str, invoice_id: str) -> dict[str, object]: """Check the status of a Lightning Network top-up invoice.""" from ..upstream.helpers import _instantiate_provider from ..upstream.ppqai import PPQAIUpstreamProvider 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 = await _get_upstream_provider_by_ref(session, provider_id) # For Routstr providers, proxy the status check if provider.provider_type == "routstr": @@ -1205,14 +1301,12 @@ async def check_topup_status(provider_id: int, invoice_id: str) -> dict[str, obj "/api/upstream-providers/{provider_id}/balance", dependencies=[Depends(require_admin_api)], ) -async def get_provider_balance(provider_id: int) -> dict[str, object]: +async def get_provider_balance(provider_id: str) -> dict[str, object]: """Get the current balance for an upstream provider account.""" from ..upstream.helpers import _instantiate_provider 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 = await _get_upstream_provider_by_ref(session, provider_id) # For Routstr providers, proxy the balance check if provider.provider_type == "routstr": @@ -1591,15 +1685,13 @@ async def get_lightning_invoices_api( "/api/upstream-providers/{provider_id}/routstr/refund", dependencies=[Depends(require_admin_api)], ) -async def refund_routstr_provider_balance(provider_id: int) -> dict[str, object]: +async def refund_routstr_provider_balance(provider_id: str) -> dict[str, object]: """Refund balance from an upstream Routstr provider back to the local wallet.""" from ..upstream.helpers import _instantiate_provider from ..upstream.routstr import RoutstrUpstreamProvider async with create_session() as session: - provider_row = await session.get(UpstreamProviderRow, provider_id) - if not provider_row: - raise HTTPException(status_code=404, detail="Provider not found") + provider_row = await _get_upstream_provider_by_ref(session, provider_id) if provider_row.provider_type != "routstr": raise HTTPException( diff --git a/routstr/core/db.py b/routstr/core/db.py index 282b8b93..586f467e 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -319,6 +319,12 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore ), ) id: int | None = Field(default=None, primary_key=True) + slug: str | None = Field( + default=None, + unique=True, + index=True, + description="Stable external slug used for updates via API key.", + ) provider_type: str = Field( description="Provider type: custom, openai, anthropic, azure, openrouter, etc." ) diff --git a/routstr/core/provider_slugs.py b/routstr/core/provider_slugs.py new file mode 100644 index 00000000..0efb919f --- /dev/null +++ b/routstr/core/provider_slugs.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +import re +from itertools import count +from typing import Collection + +from sqlmodel import select +from sqlmodel.ext.asyncio.session import AsyncSession + +from .db import UpstreamProviderRow + +_SLUG_BASE_PATTERN = re.compile(r"[^a-z0-9]+") +_MAX_SLUG_LENGTH = 64 + + +def provider_slug_base(provider_type: str) -> str: + """Return a deterministic slug base for a provider type.""" + base = _SLUG_BASE_PATTERN.sub("-", provider_type.lower()).strip("-") + if not base: + base = "provider" + elif base.isdigit(): + base = f"provider-{base}" + elif len(base) < 3: + base = f"{base}-provider" + + if len(base) > _MAX_SLUG_LENGTH: + base = base[:_MAX_SLUG_LENGTH].rstrip("-") or "provider" + return base + + +def provider_slug_candidate(base: str, suffix_number: int) -> str: + if suffix_number == 1: + return base + + suffix = f"-{suffix_number}" + max_base_length = _MAX_SLUG_LENGTH - len(suffix) + return f"{base[:max_base_length].rstrip('-')}{suffix}" + + +async def allocate_unique_provider_slug( + session: AsyncSession, + provider_type: str, + reserved_slugs: Collection[str] = (), +) -> str: + """Allocate a stable, deterministic provider slug. + + The first provider of a type gets ``openai``; later collisions get + ``openai-2``, ``openai-3``, etc. ``reserved_slugs`` covers rows staged in + memory but not flushed yet, such as settings/env seeding. + """ + base = provider_slug_base(provider_type) + reserved = {slug.lower() for slug in reserved_slugs} + + for suffix_number in count(1): + candidate = provider_slug_candidate(base, suffix_number) + if candidate in reserved: + continue + + result = await session.exec( + select(UpstreamProviderRow).where(UpstreamProviderRow.slug == candidate) + ) + if result.first() is None: + return candidate + + raise RuntimeError("unreachable") diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index 9ad4d73f..a79b8dc2 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -12,6 +12,7 @@ from sqlmodel import select from ..core import get_logger from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session +from ..core.provider_slugs import allocate_unique_provider_slug from ..payment.models import Model from .base import BaseUpstreamProvider @@ -250,6 +251,7 @@ async def _seed_providers_from_settings( providers_to_add: list[UpstreamProviderRow] = [] seeded_provider_keys: set[tuple[str, str]] = set() + reserved_slugs: set[str] = set() provider_classes_by_type = { cls.provider_type: cls @@ -279,8 +281,13 @@ async def _seed_providers_from_settings( ) ) if not result.first(): + slug = await allocate_unique_provider_slug( + session, provider_type, reserved_slugs + ) + reserved_slugs.add(slug) providers_to_add.append( UpstreamProviderRow( + slug=slug, provider_type=provider_type, base_url=base_url, api_key=api_key, @@ -299,8 +306,13 @@ async def _seed_providers_from_settings( ) ) if not result.first(): + slug = await allocate_unique_provider_slug( + session, "ollama", reserved_slugs + ) + reserved_slugs.add(slug) providers_to_add.append( UpstreamProviderRow( + slug=slug, provider_type="ollama", base_url=ollama_base_url, api_key=ollama_api_key, @@ -320,8 +332,13 @@ async def _seed_providers_from_settings( ) ) if not result.first(): + slug = await allocate_unique_provider_slug( + session, "azure", reserved_slugs + ) + reserved_slugs.add(slug) providers_to_add.append( UpstreamProviderRow( + slug=slug, provider_type="azure", base_url=base_url, api_key=api_key, @@ -342,8 +359,13 @@ async def _seed_providers_from_settings( ) ) if not result.first(): + slug = await allocate_unique_provider_slug( + session, "custom", reserved_slugs + ) + reserved_slugs.add(slug) providers_to_add.append( UpstreamProviderRow( + slug=slug, provider_type="custom", base_url=base_url, api_key=api_key, @@ -356,7 +378,7 @@ async def _seed_providers_from_settings( session.add(provider) logger.info( f"Seeding {provider.provider_type} provider", # type: ignore[str-format] - extra={"base_url": provider.base_url}, + extra={"base_url": provider.base_url, "slug": provider.slug}, ) diff --git a/tests/unit/test_provider_slug_migration.py b/tests/unit/test_provider_slug_migration.py new file mode 100644 index 00000000..c44faefb --- /dev/null +++ b/tests/unit/test_provider_slug_migration.py @@ -0,0 +1,55 @@ +from __future__ import annotations + +import importlib.util +from pathlib import Path + +import sqlalchemy as sa + +_MIGRATION_PATH = ( + Path(__file__).resolve().parents[2] + / "migrations" + / "versions" + / "c6d7e8f9a0b1_add_slug_to_upstream_providers.py" +) +_spec = importlib.util.spec_from_file_location("provider_slug_migration", _MIGRATION_PATH) +assert _spec is not None and _spec.loader is not None +migration = importlib.util.module_from_spec(_spec) +_spec.loader.exec_module(migration) + + +def test_slug_migration_backfill_uses_api_safe_deterministic_slugs() -> None: + engine = sa.create_engine("sqlite:///:memory:") + with engine.begin() as conn: + conn.execute( + sa.text( + "CREATE TABLE upstream_providers (" + "id INTEGER PRIMARY KEY, " + "provider_type VARCHAR NOT NULL, " + "slug VARCHAR NULL" + ")" + ) + ) + conn.execute( + sa.text( + "INSERT INTO upstream_providers (id, provider_type, slug) VALUES " + "(1, 'OpenAI Compatible', NULL), " + "(2, 'OpenAI Compatible', ''), " + "(3, '123', NULL), " + "(4, 'x', NULL), " + "(5, 'anthropic', 'anthropic')" + ) + ) + + migration._backfill_provider_slugs(conn) + + rows = conn.execute( + sa.text("SELECT id, slug FROM upstream_providers ORDER BY id") + ).all() + + assert rows == [ + (1, "openai-compatible"), + (2, "openai-compatible-2"), + (3, "provider-123"), + (4, "x-provider"), + (5, "anthropic"), + ] diff --git a/tests/unit/test_provider_slugs.py b/tests/unit/test_provider_slugs.py new file mode 100644 index 00000000..0d446fc1 --- /dev/null +++ b/tests/unit/test_provider_slugs.py @@ -0,0 +1,144 @@ +from __future__ import annotations + +import pytest +from sqlalchemy.ext.asyncio import create_async_engine +from sqlmodel import SQLModel, select +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import _get_upstream_provider_by_ref +from routstr.core.db import UpstreamProviderRow +from routstr.core.provider_slugs import ( + allocate_unique_provider_slug, + provider_slug_base, +) +from routstr.upstream.helpers import _seed_providers_from_settings + + +@pytest.mark.asyncio +async def test_allocate_unique_provider_slug_is_deterministic_with_suffixes() -> None: + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + async with AsyncSession(engine) as session: + session.add( + UpstreamProviderRow( + slug="openai", + provider_type="openai", + base_url="https://api.openai.com/v1", + api_key="key-1", + ) + ) + await session.commit() + + assert await allocate_unique_provider_slug(session, "openai") == "openai-2" + assert ( + await allocate_unique_provider_slug(session, "openai", {"openai-2"}) + == "openai-3" + ) + + await engine.dispose() + + +def test_provider_slug_base_sanitizes_provider_type() -> None: + assert provider_slug_base("OpenAI Compatible") == "openai-compatible" + assert provider_slug_base("!!!") == "provider" + assert provider_slug_base("AI") == "ai-provider" + assert provider_slug_base("123") == "provider-123" + + +@pytest.mark.asyncio +async def test_provider_ref_lookup_accepts_existing_numeric_ids_and_slugs() -> None: + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + async with AsyncSession(engine) as session: + provider = UpstreamProviderRow( + slug="openai", + provider_type="openai", + base_url="https://api.openai.com/v1", + api_key="key-1", + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + + by_id = await _get_upstream_provider_by_ref(session, str(provider.id)) + by_slug = await _get_upstream_provider_by_ref(session, "openai") + + assert by_id.id == provider.id + assert by_slug.id == provider.id + + await engine.dispose() + + +@pytest.mark.asyncio +async def test_seed_providers_from_settings_sets_deterministic_slug( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + monkeypatch.setenv("OPENAI_API_KEY", "seeded-openai-key") + + class SettingsStub: + chat_completions_api_version: str | None = None + upstream_base_url: str | None = None + upstream_api_key: str = "" + + async with AsyncSession(engine) as session: + await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type] + await session.commit() + + result = await session.exec(select(UpstreamProviderRow)) + providers: list[UpstreamProviderRow] = list(result.all()) + + assert [(p.provider_type, p.slug) for p in providers] == [("openai", "openai")] + + await engine.dispose() + + +@pytest.mark.asyncio +async def test_seed_providers_from_settings_keeps_slug_stable_on_reseed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + monkeypatch.setenv("OPENAI_API_KEY", "seeded-openai-key") + + class SettingsStub: + chat_completions_api_version: str | None = None + upstream_base_url: str | None = None + upstream_api_key: str = "" + + async with AsyncSession(engine) as session: + session.add( + UpstreamProviderRow( + slug="openai", + provider_type="openai", + base_url="https://example.invalid/v1", + api_key="other-key", + ) + ) + await session.commit() + + await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type] + await session.commit() + await _seed_providers_from_settings(session, SettingsStub()) # type: ignore[arg-type] + await session.commit() + + result = await session.exec( + select(UpstreamProviderRow).order_by(UpstreamProviderRow.slug) + ) + providers: list[UpstreamProviderRow] = list(result.all()) + + assert [(p.provider_type, p.slug) for p in providers] == [ + ("openai", "openai"), + ("openai", "openai-2"), + ] + + await engine.dispose() diff --git a/ui/components/provider-form-fields.tsx b/ui/components/provider-form-fields.tsx index 46a7d2ef..16e6e523 100644 --- a/ui/components/provider-form-fields.tsx +++ b/ui/components/provider-form-fields.tsx @@ -118,6 +118,26 @@ export function ProviderFormFields({ /> )} +
+ Stable external key used to update this provider via the admin API. +
+