From 770565601698d2c5df434f0f641dcea8673b0fb0 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 1 Jul 2026 16:40:31 +0200 Subject: [PATCH] resolve review comment --- routstr/core/admin.py | 21 +----- routstr/core/provider_slugs.py | 65 +++++++++++++++++ routstr/upstream/helpers.py | 24 +++++- tests/unit/test_provider_slugs.py | 117 ++++++++++++++++++++++++++++++ 4 files changed, 207 insertions(+), 20 deletions(-) create mode 100644 routstr/core/provider_slugs.py create mode 100644 tests/unit/test_provider_slugs.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 6e42b511..b0fc080e 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -31,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__) @@ -726,11 +727,6 @@ def _validate_slug(value: str) -> str: return candidate -def _generate_slug(provider_type: str) -> str: - base = re.sub(r"[^a-z0-9]+", "-", provider_type.lower()).strip("-") or "provider" - return f"{base}-{secrets.token_hex(3)}" - - async def _ensure_unique_slug( session: AsyncSession, slug: str, exclude_id: int | None = None ) -> None: @@ -883,20 +879,7 @@ async def create_upstream_provider( slug = _validate_slug(payload.slug) await _ensure_unique_slug(session, slug) else: - for _ in range(8): - slug = _generate_slug(payload.provider_type) - existing = await session.exec( - select(UpstreamProviderRow).where( - UpstreamProviderRow.slug == slug - ) - ) - if existing.first() is None: - break - else: - raise HTTPException( - status_code=500, - detail="Could not generate a unique slug", - ) + slug = await allocate_unique_provider_slug(session, payload.provider_type) provider = UpstreamProviderRow( slug=slug, diff --git a/routstr/core/provider_slugs.py b/routstr/core/provider_slugs.py new file mode 100644 index 00000000..cd0d0410 --- /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 _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 = _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_slugs.py b/tests/unit/test_provider_slugs.py new file mode 100644 index 00000000..e33fbeff --- /dev/null +++ b/tests/unit/test_provider_slugs.py @@ -0,0 +1,117 @@ +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.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_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()