mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
resolve review comment
This commit is contained in:
+2
-19
@@ -31,6 +31,7 @@ from .db import (
|
|||||||
)
|
)
|
||||||
from .log_manager import log_manager
|
from .log_manager import log_manager
|
||||||
from .logging import get_logger
|
from .logging import get_logger
|
||||||
|
from .provider_slugs import allocate_unique_provider_slug
|
||||||
from .settings import SettingsService, settings
|
from .settings import SettingsService, settings
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
@@ -726,11 +727,6 @@ def _validate_slug(value: str) -> str:
|
|||||||
return candidate
|
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(
|
async def _ensure_unique_slug(
|
||||||
session: AsyncSession, slug: str, exclude_id: int | None = None
|
session: AsyncSession, slug: str, exclude_id: int | None = None
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -883,20 +879,7 @@ async def create_upstream_provider(
|
|||||||
slug = _validate_slug(payload.slug)
|
slug = _validate_slug(payload.slug)
|
||||||
await _ensure_unique_slug(session, slug)
|
await _ensure_unique_slug(session, slug)
|
||||||
else:
|
else:
|
||||||
for _ in range(8):
|
slug = await allocate_unique_provider_slug(session, payload.provider_type)
|
||||||
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",
|
|
||||||
)
|
|
||||||
|
|
||||||
provider = UpstreamProviderRow(
|
provider = UpstreamProviderRow(
|
||||||
slug=slug,
|
slug=slug,
|
||||||
|
|||||||
@@ -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")
|
||||||
@@ -12,6 +12,7 @@ from sqlmodel import select
|
|||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session
|
from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session
|
||||||
|
from ..core.provider_slugs import allocate_unique_provider_slug
|
||||||
from ..payment.models import Model
|
from ..payment.models import Model
|
||||||
from .base import BaseUpstreamProvider
|
from .base import BaseUpstreamProvider
|
||||||
|
|
||||||
@@ -250,6 +251,7 @@ async def _seed_providers_from_settings(
|
|||||||
|
|
||||||
providers_to_add: list[UpstreamProviderRow] = []
|
providers_to_add: list[UpstreamProviderRow] = []
|
||||||
seeded_provider_keys: set[tuple[str, str]] = set()
|
seeded_provider_keys: set[tuple[str, str]] = set()
|
||||||
|
reserved_slugs: set[str] = set()
|
||||||
|
|
||||||
provider_classes_by_type = {
|
provider_classes_by_type = {
|
||||||
cls.provider_type: cls
|
cls.provider_type: cls
|
||||||
@@ -279,8 +281,13 @@ async def _seed_providers_from_settings(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if not result.first():
|
if not result.first():
|
||||||
|
slug = await allocate_unique_provider_slug(
|
||||||
|
session, provider_type, reserved_slugs
|
||||||
|
)
|
||||||
|
reserved_slugs.add(slug)
|
||||||
providers_to_add.append(
|
providers_to_add.append(
|
||||||
UpstreamProviderRow(
|
UpstreamProviderRow(
|
||||||
|
slug=slug,
|
||||||
provider_type=provider_type,
|
provider_type=provider_type,
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
@@ -299,8 +306,13 @@ async def _seed_providers_from_settings(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if not result.first():
|
if not result.first():
|
||||||
|
slug = await allocate_unique_provider_slug(
|
||||||
|
session, "ollama", reserved_slugs
|
||||||
|
)
|
||||||
|
reserved_slugs.add(slug)
|
||||||
providers_to_add.append(
|
providers_to_add.append(
|
||||||
UpstreamProviderRow(
|
UpstreamProviderRow(
|
||||||
|
slug=slug,
|
||||||
provider_type="ollama",
|
provider_type="ollama",
|
||||||
base_url=ollama_base_url,
|
base_url=ollama_base_url,
|
||||||
api_key=ollama_api_key,
|
api_key=ollama_api_key,
|
||||||
@@ -320,8 +332,13 @@ async def _seed_providers_from_settings(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if not result.first():
|
if not result.first():
|
||||||
|
slug = await allocate_unique_provider_slug(
|
||||||
|
session, "azure", reserved_slugs
|
||||||
|
)
|
||||||
|
reserved_slugs.add(slug)
|
||||||
providers_to_add.append(
|
providers_to_add.append(
|
||||||
UpstreamProviderRow(
|
UpstreamProviderRow(
|
||||||
|
slug=slug,
|
||||||
provider_type="azure",
|
provider_type="azure",
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
@@ -342,8 +359,13 @@ async def _seed_providers_from_settings(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if not result.first():
|
if not result.first():
|
||||||
|
slug = await allocate_unique_provider_slug(
|
||||||
|
session, "custom", reserved_slugs
|
||||||
|
)
|
||||||
|
reserved_slugs.add(slug)
|
||||||
providers_to_add.append(
|
providers_to_add.append(
|
||||||
UpstreamProviderRow(
|
UpstreamProviderRow(
|
||||||
|
slug=slug,
|
||||||
provider_type="custom",
|
provider_type="custom",
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
@@ -356,7 +378,7 @@ async def _seed_providers_from_settings(
|
|||||||
session.add(provider)
|
session.add(provider)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Seeding {provider.provider_type} provider", # type: ignore[str-format]
|
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},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user